Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
561ba3409f | ||
|
|
6cf67e98aa | ||
|
|
b324f132d8 | ||
|
|
c84578d9f9 | ||
|
|
ae2dbc7e1a | ||
|
|
8e553f5662 | ||
|
|
375d603a8e | ||
|
|
17c118cba2 | ||
|
|
51f4b8464e |
@@ -61,8 +61,8 @@ If a required tool is missing from `..\..\venv`, install or update development d
|
||||
## Architecture Rules
|
||||
|
||||
- Organize code into clear layers with one-way dependencies.
|
||||
- ComfyUI integration layer: `NODE_CLASS_MAPPINGS`, `NODE_DISPLAY_NAME_MAPPINGS`, node categories, input/output declarations, and ComfyUI import-time registration.
|
||||
- Node API layer: thin node classes exposing ComfyUI-facing methods such as `INPUT_TYPES`, `RETURN_TYPES`, `FUNCTION`, and execution entry points.
|
||||
- ComfyUI integration layer: the Comfy v3 `comfy_entrypoint()`, `simple_syrup/nodes_v3/__init__.py::get_nodes()`, v3 schema declarations, node categories, input/output declarations, and ComfyUI import-time registration.
|
||||
- Node API layer: thin v3 node classes exposing ComfyUI-facing schema and execution entry points.
|
||||
- Application/service layer: orchestration for node behavior, validation flow, and feature-level use cases.
|
||||
- Domain layer: stable internal models, value objects, policies, and pure behavior.
|
||||
- Runtime/adapter layer: filesystem access, image/audio/model IO, ComfyUI object adaptation, subprocess boundaries, network boundaries, and optional external integrations.
|
||||
@@ -83,7 +83,7 @@ If a required tool is missing from `..\..\venv`, install or update development d
|
||||
- For behavior-critical areas, work in two steps:
|
||||
1. Add characterization/regression tests for existing behavior.
|
||||
2. Perform structural changes behind those tests.
|
||||
- Behavior-critical areas include node registration, `INPUT_TYPES`, `RETURN_TYPES`, widget names, output ordering, execution return shapes, workflow compatibility, validation behavior, file IO, model IO, image/audio tensor handling, and ComfyUI import behavior.
|
||||
- Behavior-critical areas include node registration, v3 schema declarations, return metadata, widget names, output ordering, execution return shapes, workflow compatibility, validation behavior, file IO, model IO, image/audio tensor handling, and ComfyUI import behavior.
|
||||
- Do not start structural changes in an area without behavior safeguards for that area.
|
||||
- When behavior spans multiple components, trace the current ownership and data flow before editing.
|
||||
- Correct the ownership model instead of layering compensating patches across consumers.
|
||||
@@ -99,7 +99,7 @@ If a required tool is missing from `..\..\venv`, install or update development d
|
||||
- Public node identifiers are compatibility-sensitive.
|
||||
- Do not rename node classes, display names, categories, input keys, output names, return types, or function names without explicit approval.
|
||||
- Keep ComfyUI-facing node classes small and predictable.
|
||||
- `INPUT_TYPES` must be deterministic and must not perform expensive IO.
|
||||
- V3 schema declarations must be deterministic and must not perform expensive IO.
|
||||
- Importing the node pack must not perform heavy computation, network access, model loading, or destructive filesystem operations.
|
||||
- Node execution must validate inputs before performing side effects.
|
||||
- Node execution must return exactly the declared output shape.
|
||||
@@ -121,25 +121,23 @@ If a required tool is missing from `..\..\venv`, install or update development d
|
||||
- Avoid implementation jargon unless the user needs it to make a good workflow decision.
|
||||
- Do not repeat the field name as a definition.
|
||||
- Do not document removed behavior, imagined alternatives, or choices the product does not expose.
|
||||
- Keep legacy `INPUT_TYPES` tooltips and Comfy v3 schema tooltips aligned when both export paths expose the same node or field.
|
||||
- Keep Comfy v3 schema tooltips aligned with the current node behavior.
|
||||
|
||||
## ComfyUI Node Export Rules
|
||||
|
||||
- When adding, renaming, or removing a ComfyUI node, update and verify every export path used by this repository.
|
||||
- Legacy ComfyUI mapping exports must be updated in `simple_syrup/nodes/__init__.py`:
|
||||
- `NODE_CLASS_MAPPINGS`
|
||||
- `NODE_DISPLAY_NAME_MAPPINGS`
|
||||
- `__all__`
|
||||
- Comfy v3 entrypoint exports must be updated when the node should be visible through the v3 API:
|
||||
- Comfy v3 is the only supported ComfyUI node export path.
|
||||
- Do not add, preserve, or restore legacy ComfyUI mapping exports.
|
||||
- The root package export in repository root `__init__.py` must expose `comfy_entrypoint` and must not expose `NODE_CLASS_MAPPINGS` or `NODE_DISPLAY_NAME_MAPPINGS`.
|
||||
- `simple_syrup/nodes_v3/__init__.py::get_nodes()` is the authoritative node registry.
|
||||
- When adding, renaming, or removing a ComfyUI node, update and verify:
|
||||
- the v3 wrapper class when needed
|
||||
- `simple_syrup/nodes_v3/__init__.py`
|
||||
- `get_nodes()`
|
||||
- a v3 wrapper class when needed
|
||||
- The root package export in repository root `__init__.py` must continue exposing the relevant mappings and `comfy_entrypoint`.
|
||||
- Tests must cover every export path used by the node:
|
||||
- A registration test must assert the node id and display name exist in `NODE_CLASS_MAPPINGS` and `NODE_DISPLAY_NAME_MAPPINGS`.
|
||||
- A v3 entrypoint test must assert `comfy_entrypoint().get_node_list()` includes the node when it is expected to be visible through Comfy v3.
|
||||
- If a v3 node is conditional, tests must cover both the available and unavailable conditions and prove unrelated v3 nodes remain exported.
|
||||
- Do not consider a node addition complete from `NODE_CLASS_MAPPINGS` alone. A node is not fully exported until every repository-supported ComfyUI export path is updated and tested.
|
||||
- tests that exercise the v3 schema, entrypoint export, and execution behavior
|
||||
- Maintained nodes must keep stable `SimpleSyrup.*` node ids unless the maintainer explicitly approves a workflow-facing rename.
|
||||
- If a v3 node is conditional, tests must cover both the available and unavailable conditions and prove unrelated v3 nodes remain exported.
|
||||
- Do not maintain a compatibility layer, fallback registry, or migration path for old ComfyUI versions that only load `NODE_CLASS_MAPPINGS`.
|
||||
- A node is not fully exported until the v3 entrypoint returns it and tests prove the intended schema and behavior.
|
||||
|
||||
## Code Organization and Readability
|
||||
|
||||
@@ -215,7 +213,7 @@ If a required tool is missing from `..\..\venv`, install or update development d
|
||||
- Use real behavior tests over excessive mocking.
|
||||
- Mock only external boundaries such as ComfyUI runtime calls, filesystem errors, network calls, subprocesses, random generation, and time.
|
||||
- Node behavior must be tested at the narrowest useful level and through integration-style tests when ComfyUI-facing shape matters.
|
||||
- Node registration changes require tests for exported mappings.
|
||||
- Node registration changes require tests for the v3 entrypoint and v3 schemas.
|
||||
- Tooltip coverage must be tested for exported node descriptions, inputs, and outputs supported by each ComfyUI API path.
|
||||
- Input/output signature changes require workflow-facing compatibility tests.
|
||||
- Runtime behavior requires tests for success and failure paths.
|
||||
@@ -301,7 +299,7 @@ Per change, all of the following are required:
|
||||
- Frontend source is typed and tested when touched.
|
||||
- Generated frontend artifacts are rebuilt from source.
|
||||
- `npm run lint:web`, `npm run typecheck:web`, `npm run test:web`, and `npm run build:web` pass when frontend code exists or is touched.
|
||||
- New, renamed, or removed nodes are updated in all legacy and Comfy v3 export paths, with tests proving both paths expose the intended node set.
|
||||
- New, renamed, or removed nodes are updated in the Comfy v3 export path, with tests proving the v3 entrypoint exposes the intended node set.
|
||||
|
||||
## Commit Policy
|
||||
|
||||
|
||||
@@ -1,3 +1,30 @@
|
||||
# [1.4.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.3.0...v1.4.0) (2026-06-02)
|
||||
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* **tiled-diffusion:** clamp overlap for small latents ([d7448c6](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/d7448c6ca52ce517b8d0f8ee697249c5def13535))
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **detailing:** add external llm segs tagging ([b7cd40c](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/b7cd40c85ae9f9514d0b1ff3c3f6007730796752))
|
||||
* **segs:** add regional batching and wd14 tagging nodes ([fef60ad](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/fef60adeca2c8ec7c0641e8106bb0863ee2f195e))
|
||||
|
||||
# [1.3.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.2.0...v1.3.0) (2026-05-26)
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **prompt-control:** add schedule and encode prompt node ([bd515e6](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/bd515e696cedcc78d056af6c23b9193e34f131bc))
|
||||
|
||||
# [1.2.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.1.0...v1.2.0) (2026-05-25)
|
||||
|
||||
|
||||
### Features
|
||||
|
||||
* **nodes:** add VAE options and clone-safe diffusion ([6ff2dc8](https://github.com/Artificial-Sweetener/SimpleSyrup/commit/6ff2dc8c24a6f7ddde3182b81bcbe6aad65427f4))
|
||||
|
||||
# [1.1.0](https://github.com/Artificial-Sweetener/SimpleSyrup/compare/v1.0.0...v1.1.0) (2026-05-23)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,670 @@
|
||||
# Plan: Tag SEGS w/ External LLM
|
||||
|
||||
## Goal
|
||||
|
||||
Add a new Comfy v3-only node named `Tag SEGS w/ External LLM`.
|
||||
|
||||
The node will take an image, existing Impact-compatible SEGS, and a CLIP object. It will show each SEG crop to the configured external vision LLM, turn each LLM response into a positive regional prompt, CLIP-encode those prompts in the same order as the SEGS, and return the unchanged SEGS plus an aligned `CONDITIONING_BATCH` output named `positive`.
|
||||
|
||||
The intended workflow is the same regional-conditioning role currently served by `Tag SEGS w/ WD14` and `Tile & Tag SEGS`, but with an external vision model instead of WD14. The node should be useful when a vision LLM can produce better region descriptions than WD14 tags.
|
||||
|
||||
## Decisions From Planning
|
||||
|
||||
- Export only through the Comfy v3 node path.
|
||||
- Do not add a legacy node class.
|
||||
- Do not add `NODE_CLASS_MAPPINGS` or `NODE_DISPLAY_NAME_MAPPINGS`.
|
||||
- Do not add an internal compatibility shim or dual implementation path.
|
||||
- Make exactly one external LLM call per SEG.
|
||||
- Preserve SEGS order exactly.
|
||||
- Return the original SEGS unchanged, converted only to the existing Impact-compatible output shape.
|
||||
- Output one `CONDITIONING_BATCH` entry per SEG.
|
||||
- Fail closed when alignment is uncertain.
|
||||
- Provide a single widget for how SEG crops are presented to the vision model.
|
||||
- Support LLM response cleanup controls similar to the WD14 tagging nodes.
|
||||
- Treat the LLM response as prompt text after lightweight tag formatting. Do not add clever semantic parsing in the first implementation.
|
||||
- Do not provide opinionated default prompt text for `system_prompt` or `user_prompt`; leave both defaults empty.
|
||||
|
||||
## Implementation Status
|
||||
|
||||
- [x] Review plan and current code patterns.
|
||||
- Landing note: existing `Tag SEGS w/ WD14` owns SEGS validation/alignment, `External LLM Prompt` owns provider settings/key flow, and v3 registration is centralized in `simple_syrup/nodes_v3/__init__.py`.
|
||||
- [x] Add SEG crop image encoder and tests.
|
||||
- Landing note: `ExternalLLMSegsImageEncoder` now emits PNG data URLs for `transparent mask`, `black mask`, and `full crop`; focused encoder tests pass.
|
||||
- [x] Add external LLM image-data-url execution path.
|
||||
- Landing note: `ExternalLLMPromptService.generate_with_image_data_url()` reuses the existing model/settings/key/client flow while accepting the SEG encoder's prebuilt PNG data URL.
|
||||
- [x] Add service formatting and orchestration tests.
|
||||
- Landing note: service tests cover per-SEG LLM call order, prompt formatting, exclusions, progress, logging metadata, validation failures, and conditioning count alignment.
|
||||
- [x] Add `TagSEGSWithExternalLLMService`.
|
||||
- Landing note: the service owns SEGS validation, SEG image encoding, pre-encoded LLM calls, response formatting, prompt prefixing, CLIP conditioning encoding, and completion logging.
|
||||
- [x] Add v3-only node schema and execution forwarding.
|
||||
- Landing note: `TagSEGSWithExternalLLMV3` defines the schema directly and delegates to `TagSEGSWithExternalLLMService`; no file was added under `simple_syrup/nodes/`. `system_prompt` and `user_prompt` default to empty strings.
|
||||
- [x] Register v3 node and update registration tests.
|
||||
- Landing note: `SimpleSyrup.TagSEGSWithExternalLLM` is now returned by `get_nodes()` with the maintained base node list.
|
||||
- [x] Add tooltips and update tooltip coverage.
|
||||
- Landing note: all new v3 inputs/outputs use dedicated tooltip constants, and focused registration/tooltip tests pass.
|
||||
- [x] Update this plan with final landing notes.
|
||||
- Landing note: implementation completed as a v3-only vertical slice with no new `simple_syrup/nodes/` node class and no legacy mapping export changes.
|
||||
- [x] Run required verification gates.
|
||||
- Landing note: `ruff format .`, `ruff check .`, `mypy --strict simple_syrup tests`, and `pytest -n auto -q` pass in `..\..\venv`.
|
||||
|
||||
## New Node Contract
|
||||
|
||||
Node id:
|
||||
|
||||
```text
|
||||
SimpleSyrup.TagSEGSWithExternalLLM
|
||||
```
|
||||
|
||||
Display name:
|
||||
|
||||
```text
|
||||
Tag SEGS w/ External LLM
|
||||
```
|
||||
|
||||
Category:
|
||||
|
||||
```text
|
||||
SimpleSyrup/Detailing
|
||||
```
|
||||
|
||||
Inputs:
|
||||
|
||||
```text
|
||||
image: IMAGE
|
||||
segs: SEGS
|
||||
clip: CLIP
|
||||
model: external LLM model dropdown
|
||||
system_prompt: STRING
|
||||
user_prompt: STRING
|
||||
universal_positive: STRING
|
||||
seg_image_mode: COMBO
|
||||
replace_underscore: BOOLEAN
|
||||
trailing_comma: BOOLEAN
|
||||
exclude_tags: STRING
|
||||
max_tokens: INT
|
||||
reasoning_effort: COMBO
|
||||
```
|
||||
|
||||
Outputs:
|
||||
|
||||
```text
|
||||
segs: SEGS
|
||||
positive: CONDITIONING_BATCH
|
||||
```
|
||||
|
||||
Suggested default input values:
|
||||
|
||||
```text
|
||||
system_prompt:
|
||||
|
||||
user_prompt:
|
||||
|
||||
universal_positive:
|
||||
|
||||
seg_image_mode:
|
||||
transparent mask
|
||||
|
||||
replace_underscore:
|
||||
true
|
||||
|
||||
trailing_comma:
|
||||
false
|
||||
|
||||
exclude_tags:
|
||||
|
||||
max_tokens:
|
||||
1024
|
||||
|
||||
reasoning_effort:
|
||||
default
|
||||
```
|
||||
|
||||
Search aliases:
|
||||
|
||||
```text
|
||||
llm
|
||||
vision
|
||||
tag
|
||||
segs
|
||||
detail
|
||||
regional
|
||||
prompt
|
||||
```
|
||||
|
||||
## SEG Image Modes
|
||||
|
||||
Add one combo input named `seg_image_mode` with these exact options:
|
||||
|
||||
```text
|
||||
transparent mask
|
||||
black mask
|
||||
full crop
|
||||
```
|
||||
|
||||
### `transparent mask`
|
||||
|
||||
Crop the original image to `segment.crop_region`.
|
||||
|
||||
Crop the SEG mask to the same rectangle. Resize or normalize the mask only if needed to match the crop dimensions.
|
||||
|
||||
Encode a PNG data URL with RGBA channels:
|
||||
|
||||
- RGB comes from the crop.
|
||||
- Alpha comes from the cropped SEG mask.
|
||||
- Pixels outside the mask are hidden by alpha.
|
||||
|
||||
Use this as the default because it most directly communicates "only this SEG matters" when the provider supports alpha.
|
||||
|
||||
### `black mask`
|
||||
|
||||
Crop the original image to `segment.crop_region`.
|
||||
|
||||
Crop the SEG mask to the same rectangle. Resize or normalize the mask only if needed to match the crop dimensions.
|
||||
|
||||
Encode an RGB PNG data URL:
|
||||
|
||||
- Pixels inside the mask keep the original crop color.
|
||||
- Pixels outside the mask are set to black.
|
||||
|
||||
This is the compatibility mode for providers that ignore PNG alpha or flatten transparent pixels unpredictably.
|
||||
|
||||
### `full crop`
|
||||
|
||||
Crop the original image to `segment.crop_region`.
|
||||
|
||||
Encode the full RGB crop unchanged.
|
||||
|
||||
Ignore the SEG mask for the image sent to the LLM. This gives the model surrounding context when context improves region tagging.
|
||||
|
||||
## Formatting Rules For LLM Responses
|
||||
|
||||
Add formatting controls:
|
||||
|
||||
```text
|
||||
replace_underscore: BOOLEAN = true
|
||||
trailing_comma: BOOLEAN = false
|
||||
exclude_tags: STRING = ""
|
||||
```
|
||||
|
||||
For each LLM response:
|
||||
|
||||
1. Strip leading and trailing whitespace.
|
||||
2. Reject the response if it is empty after stripping.
|
||||
3. Split the response by commas into tag-like chunks.
|
||||
4. Trim whitespace around each chunk.
|
||||
5. Drop empty chunks.
|
||||
6. If `replace_underscore` is enabled, replace `_` with spaces inside each chunk.
|
||||
7. Parse `exclude_tags` as a comma-separated list.
|
||||
8. Normalize excluded tags using the same trim, empty-drop, and optional underscore replacement policy.
|
||||
9. Remove exact tag matches after normalization.
|
||||
10. Rejoin remaining tags with `, `.
|
||||
11. Reject the prompt if no tags remain after exclusion.
|
||||
12. If `trailing_comma` is enabled and the prompt does not already end in a comma, append `,`.
|
||||
13. Prefix with `universal_positive` using `domain.prompt_composition.prefix_prompt`.
|
||||
|
||||
Do not fuzzy-match excluded tags.
|
||||
|
||||
Do not parse JSON, markdown, prose sections, or model-specific response schemas in the first implementation. Prompt discipline should keep responses comma-like. If a model returns prose, the comma splitting still gives deterministic behavior.
|
||||
|
||||
## Architecture
|
||||
|
||||
Follow the existing layer boundaries in `AGENTS.md`.
|
||||
|
||||
### Domain Layer
|
||||
|
||||
Use existing domain objects where possible:
|
||||
|
||||
- `simple_syrup/domain/segs.py`
|
||||
- `simple_syrup/domain/conditioning_batch.py`
|
||||
- `simple_syrup/domain/prompt_composition.py`
|
||||
- `simple_syrup/domain/external_llm.py`
|
||||
|
||||
Add a small domain value object only if it removes duplication or clarifies validation. A likely useful object is a frozen dataclass for formatting controls:
|
||||
|
||||
```python
|
||||
@dataclass(frozen=True)
|
||||
class LLMTagFormattingControls:
|
||||
replace_underscore: bool
|
||||
trailing_comma: bool
|
||||
exclude_tags: str
|
||||
```
|
||||
|
||||
Place it where ownership is clearest after implementation. If the formatting is used only by this service, keep it service-local. If it becomes shared with future LLM tagging code, place it in a domain module.
|
||||
|
||||
### Runtime / Adapter Layer
|
||||
|
||||
Add a crop image encoder near:
|
||||
|
||||
```text
|
||||
simple_syrup/runtime/external_llm_images.py
|
||||
```
|
||||
|
||||
Recommended additions:
|
||||
|
||||
```python
|
||||
SEG_IMAGE_MODES = ("transparent mask", "black mask", "full crop")
|
||||
|
||||
class ExternalLLMSegsImageEncoder:
|
||||
"""Encode SEG crops for OpenAI-compatible vision payloads."""
|
||||
|
||||
def encode_segment_as_data_url(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
segment: Segment,
|
||||
mode: str,
|
||||
) -> str:
|
||||
"""Return one SEG crop as a PNG data URL."""
|
||||
```
|
||||
|
||||
This adapter owns tensor-to-PIL conversion and PNG/base64 data URL encoding.
|
||||
|
||||
The existing `ExternalLLMImageEncoder.encode_first_image_as_data_url()` sends the first whole image. Do not overload that method with SEG-specific behavior. Keep the SEG crop behavior separate so the existing `External LLM Prompt` node remains simple and unchanged.
|
||||
|
||||
Use existing masking helpers where possible:
|
||||
|
||||
- `validate_single_image`
|
||||
- `crop_image`
|
||||
- `crop_mask`
|
||||
- `resize_mask`
|
||||
|
||||
If `segment.cropped_mask` is already crop-local, do not assume it is full-image. The service or encoder must normalize the mask robustly:
|
||||
|
||||
- If the mask shape matches the crop height and width, use it directly.
|
||||
- If the mask shape matches the full source image height and width, crop it by `segment.crop_region`.
|
||||
- If the mask has a batch dimension, normalize to one HW mask.
|
||||
- If the mask shape is compatible but not exact, resize with bilinear interpolation and clamp to `0.0..1.0`.
|
||||
- If the mask cannot be interpreted as HW or BHW, raise a clear `ValueError`.
|
||||
|
||||
Preserve security boundaries:
|
||||
|
||||
- No filesystem writes are needed.
|
||||
- No workflow-provided code execution.
|
||||
- No shell invocation.
|
||||
- No network logic inside the image encoder.
|
||||
|
||||
### Application / Service Layer
|
||||
|
||||
Add a service module:
|
||||
|
||||
```text
|
||||
simple_syrup/services/tag_segs_with_external_llm_service.py
|
||||
```
|
||||
|
||||
Recommended public result:
|
||||
|
||||
```python
|
||||
@dataclass(frozen=True)
|
||||
class TagSEGSWithExternalLLMResult:
|
||||
"""Return unchanged SEGS and aligned external-LLM conditioning."""
|
||||
|
||||
segs: ImpactSegs
|
||||
positive: ConditioningBatch
|
||||
```
|
||||
|
||||
Recommended service:
|
||||
|
||||
```python
|
||||
class TagSEGSWithExternalLLMService:
|
||||
"""Caption provided SEGS crops with an external LLM and encode conditioning."""
|
||||
```
|
||||
|
||||
The service should own:
|
||||
|
||||
- Input validation.
|
||||
- External LLM model resolution.
|
||||
- External LLM settings and credential checks.
|
||||
- One provider request per SEG.
|
||||
- Progress reporting.
|
||||
- Response formatting.
|
||||
- Prompt prefixing.
|
||||
- CLIP conditioning encoding.
|
||||
- Alignment checks.
|
||||
- Structured logging.
|
||||
|
||||
Keep external boundaries injectable for tests:
|
||||
|
||||
```python
|
||||
class VisionLLMBoundary(Protocol):
|
||||
def generate(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int,
|
||||
reasoning_effort: str,
|
||||
image: object | None = None,
|
||||
) -> str:
|
||||
"""Return one assistant response."""
|
||||
|
||||
|
||||
class SegmentImageEncodingBoundary(Protocol):
|
||||
def encode_segment_as_data_url(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
segment: Segment,
|
||||
mode: str,
|
||||
) -> str:
|
||||
"""Return one SEG crop image data URL."""
|
||||
|
||||
|
||||
class ConditioningEncodingBoundary(Protocol):
|
||||
def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch:
|
||||
"""Return conditioning entries in prompt order."""
|
||||
```
|
||||
|
||||
The existing `ExternalLLMPromptService.generate()` expects an `image` object and then asks its own image encoder to encode it. For this node, prefer adding a provider-call method that accepts a prebuilt image data URL, or extract shared provider execution into a helper. Do not pass already encoded data URLs through an API that claims to accept Comfy `IMAGE` tensors.
|
||||
|
||||
Recommended clean shape:
|
||||
|
||||
- Keep `ExternalLLMPromptService.generate()` behavior unchanged.
|
||||
- Add a method to the same service or a new small collaborator that accepts `image_data_url`.
|
||||
- Reuse the same settings repository, key store, model choice, and provider client logic.
|
||||
- Avoid duplicating endpoint/API-key validation across services if it can be extracted without creating an unclear abstraction.
|
||||
|
||||
The service execution should look like:
|
||||
|
||||
1. Validate `image` as a single BHWC IMAGE tensor.
|
||||
2. Coerce `segs` with `coerce_segs`.
|
||||
3. Validate SEGS header matches the image height and width.
|
||||
4. Validate every `segment.crop_region` fits inside the image.
|
||||
5. Reject empty SEGS.
|
||||
6. Resolve the selected LLM model.
|
||||
7. Create progress with `len(segments) + 2`.
|
||||
8. For each segment in input order:
|
||||
- Encode one SEG crop according to `seg_image_mode`.
|
||||
- Call the external LLM once with that crop.
|
||||
- Format the LLM response.
|
||||
- Prefix with `universal_positive`.
|
||||
- Append the prompt to an ordered tuple/list.
|
||||
- Update progress.
|
||||
9. Encode the prompt tuple using `ComfyConditioningEncoder.encode_batch`.
|
||||
10. Verify `len(positive.entries) == len(segments)`.
|
||||
11. Log completion with operation name, segment count, selected model, image mode, and formatting-control presence.
|
||||
12. Return `to_impact_compatible_segs(native_segs)` and `positive`.
|
||||
|
||||
Suggested operation constant:
|
||||
|
||||
```python
|
||||
OPERATION = "Tag SEGS w/ External LLM"
|
||||
```
|
||||
|
||||
### Comfy v3 Node Layer
|
||||
|
||||
Add only:
|
||||
|
||||
```text
|
||||
simple_syrup/nodes_v3/tag_segs_with_external_llm.py
|
||||
```
|
||||
|
||||
Do not add:
|
||||
|
||||
```text
|
||||
simple_syrup/nodes/tag_segs_with_external_llm.py
|
||||
```
|
||||
|
||||
The v3 node should:
|
||||
|
||||
- Define the schema directly.
|
||||
- Use `comfy_api.latest`.
|
||||
- Use `Custom("CONDITIONING_BATCH")` for the positive output.
|
||||
- Use model choices from the external LLM service, like the current `External LLM Prompt` node does.
|
||||
- Delegate non-trivial behavior to `TagSEGSWithExternalLLMService`.
|
||||
- Contain only input declaration and execution forwarding.
|
||||
|
||||
Register the node in:
|
||||
|
||||
```text
|
||||
simple_syrup/nodes_v3/__init__.py
|
||||
```
|
||||
|
||||
The root package export must continue to expose only `comfy_entrypoint`.
|
||||
|
||||
## Tooltips
|
||||
|
||||
Update:
|
||||
|
||||
```text
|
||||
simple_syrup/nodes/tooltips.py
|
||||
```
|
||||
|
||||
Every visible input and output needs a concise user-facing tooltip.
|
||||
|
||||
Suggested tooltip intent:
|
||||
|
||||
- `image`: The source image that the SEGS were detected from.
|
||||
- `segs`: The regions to describe and align with conditioning.
|
||||
- `clip`: The CLIP/text encoder used to encode generated regional prompts.
|
||||
- `model`: The configured external vision model used for each SEG crop.
|
||||
- `system_prompt`: Instructions that control how the model writes tags.
|
||||
- `user_prompt`: The per-region request sent with each SEG crop.
|
||||
- `universal_positive`: Prompt text added before every generated region prompt.
|
||||
- `seg_image_mode`: How pixels outside the SEG mask are shown to the vision model.
|
||||
- `replace_underscore`: Converts booru-style underscores into spaces before encoding.
|
||||
- `trailing_comma`: Adds a final comma to each generated prompt when enabled.
|
||||
- `exclude_tags`: Comma-separated exact tags to remove from LLM responses.
|
||||
- `max_tokens`: Maximum response length for each SEG request.
|
||||
- `reasoning_effort`: Provider-specific reasoning control, when supported.
|
||||
- `segs` output: The original SEGS, kept in the same order as the conditioning batch.
|
||||
- `positive` output: One CLIP-encoded positive conditioning entry per SEG.
|
||||
|
||||
Keep wording concise. Do not document removed choices or alternatives the node does not expose.
|
||||
|
||||
## Tests
|
||||
|
||||
Add or update tests alongside implementation.
|
||||
|
||||
### Image Encoder Tests
|
||||
|
||||
Add:
|
||||
|
||||
```text
|
||||
tests/test_external_llm_segs_images.py
|
||||
```
|
||||
|
||||
Cover:
|
||||
|
||||
- `transparent mask` returns a PNG data URL with alpha.
|
||||
- Alpha is low/zero outside the SEG mask and high/one inside it.
|
||||
- `black mask` returns RGB where outside-mask pixels are black.
|
||||
- `full crop` preserves the crop rectangle without masking.
|
||||
- Crop dimensions match `segment.crop_region`.
|
||||
- Crop-local masks are accepted.
|
||||
- Full-image masks are accepted and cropped.
|
||||
- Invalid mask rank fails with an actionable `ValueError`.
|
||||
- Unknown `seg_image_mode` fails with an actionable `ValueError`.
|
||||
|
||||
Use small deterministic tensors, such as 4x4 or 6x6 images, so pixel assertions are simple.
|
||||
|
||||
### Service Tests
|
||||
|
||||
Add:
|
||||
|
||||
```text
|
||||
tests/test_tag_segs_with_external_llm_service.py
|
||||
```
|
||||
|
||||
Cover:
|
||||
|
||||
- Existing SEGS, LLM responses, formatted prompts, and conditioning stay aligned by index.
|
||||
- The service makes one LLM call per SEG.
|
||||
- The service sends the selected `seg_image_mode` to the encoder.
|
||||
- `universal_positive` is prefixed to every generated prompt.
|
||||
- `replace_underscore=True` turns `blue_hair` into `blue hair`.
|
||||
- `replace_underscore=False` preserves `blue_hair`.
|
||||
- `exclude_tags` removes exact normalized tags.
|
||||
- `trailing_comma=True` adds a comma after formatting.
|
||||
- Empty SEGS raise `ValueError`.
|
||||
- Mismatched SEGS header and image dimensions raise `ValueError`.
|
||||
- Crop regions outside the image raise `ValueError`.
|
||||
- Empty LLM responses raise `ValueError`.
|
||||
- Responses that become empty after exclusions raise `ValueError`.
|
||||
- Conditioning encoder returning the wrong entry count raises `ValueError`.
|
||||
- Completion logging includes operation, segment count, model, image mode, and universal-positive presence.
|
||||
|
||||
Use fake boundaries for LLM calls, image encoding, conditioning encoding, and progress. Do not call a real provider in tests.
|
||||
|
||||
### V3 Node Tests
|
||||
|
||||
Add:
|
||||
|
||||
```text
|
||||
tests/test_tag_segs_with_external_llm_v3_node.py
|
||||
```
|
||||
|
||||
Cover:
|
||||
|
||||
- Schema node id is `SimpleSyrup.TagSEGSWithExternalLLM`.
|
||||
- Display name is `Tag SEGS w/ External LLM`.
|
||||
- Category is `SimpleSyrup/Detailing`.
|
||||
- Inputs are exposed in the intended order.
|
||||
- `seg_image_mode` choices are exactly `transparent mask`, `black mask`, `full crop`.
|
||||
- `transparent mask` is the default.
|
||||
- Formatting controls have intended defaults.
|
||||
- Outputs are `segs` and `positive`.
|
||||
- Output types are `SEGS` and `CONDITIONING_BATCH`.
|
||||
- `execute()` forwards all inputs to the service.
|
||||
|
||||
### Registration Tests
|
||||
|
||||
Update existing registration/v3 export tests. The new node must appear in the v3 entrypoint output from:
|
||||
|
||||
```text
|
||||
simple_syrup/nodes_v3/__init__.py::get_nodes()
|
||||
```
|
||||
|
||||
Also update tooltip coverage tests so the new node does not create coverage gaps.
|
||||
|
||||
## Validation And Errors
|
||||
|
||||
Use explicit, actionable errors.
|
||||
|
||||
Required validation failures:
|
||||
|
||||
- `image` is not a torch IMAGE tensor.
|
||||
- `image` batch size is not 1.
|
||||
- `segs` is not a valid SEGS payload.
|
||||
- SEGS header dimensions do not match the image.
|
||||
- SEGS contains no segments.
|
||||
- A SEG crop region is outside the image.
|
||||
- `seg_image_mode` is unknown.
|
||||
- The external LLM provider is not configured.
|
||||
- The external LLM API key is missing.
|
||||
- The external LLM response is empty.
|
||||
- Formatting removes every generated tag.
|
||||
- The conditioning encoder returns a count different from the SEG count.
|
||||
|
||||
Do not silently continue after invalid inputs or provider failures.
|
||||
|
||||
## Logging And Progress
|
||||
|
||||
Use `simple_syrup.shared.logging.get_logger`.
|
||||
|
||||
At completion, log an `info` event with:
|
||||
|
||||
```text
|
||||
operation: tag_segs_with_external_llm
|
||||
segment_count
|
||||
external_llm_model
|
||||
seg_image_mode
|
||||
universal_positive_present
|
||||
replace_underscore
|
||||
trailing_comma
|
||||
exclude_tags_present
|
||||
```
|
||||
|
||||
Provider failures should preserve exception context. Existing external LLM provider errors can propagate if they are already actionable.
|
||||
|
||||
Use Comfy progress in the service:
|
||||
|
||||
- Start with one initial progress update after validation/model resolution.
|
||||
- Update after each SEG LLM call.
|
||||
- Update once after conditioning encode.
|
||||
|
||||
## Security And Runtime Constraints
|
||||
|
||||
- Use the ComfyUI virtual environment located at `..\..\venv`.
|
||||
- Do not create a repository-local `.venv`.
|
||||
- Do not use global Python for verification.
|
||||
- Use PowerShell command forms on Windows.
|
||||
- Do not write temporary crop files to disk.
|
||||
- Do not log API keys, provider secrets, or full data URLs.
|
||||
- Do not execute workflow-provided code.
|
||||
- Keep network calls inside the external LLM client/provider boundary.
|
||||
- Keep filesystem, subprocess, and network behavior out of node classes and domain logic.
|
||||
- Keep import-time behavior cheap. The node schema may read cached model names, but must not make provider requests during import/schema declaration.
|
||||
|
||||
## Required Files To Touch
|
||||
|
||||
Expected new files:
|
||||
|
||||
```text
|
||||
simple_syrup/services/tag_segs_with_external_llm_service.py
|
||||
simple_syrup/nodes_v3/tag_segs_with_external_llm.py
|
||||
tests/test_external_llm_segs_images.py
|
||||
tests/test_tag_segs_with_external_llm_service.py
|
||||
tests/test_tag_segs_with_external_llm_v3_node.py
|
||||
```
|
||||
|
||||
Expected modified files:
|
||||
|
||||
```text
|
||||
simple_syrup/runtime/external_llm_images.py
|
||||
simple_syrup/services/external_llm_prompt_service.py
|
||||
simple_syrup/nodes/tooltips.py
|
||||
simple_syrup/nodes_v3/__init__.py
|
||||
tests/test_registration.py
|
||||
tests/test_node_tooltips.py
|
||||
```
|
||||
|
||||
Only modify additional files if implementation proves they are the correct ownership location.
|
||||
|
||||
Do not add a file under:
|
||||
|
||||
```text
|
||||
simple_syrup/nodes/
|
||||
```
|
||||
|
||||
for this node.
|
||||
|
||||
## Implementation Sequence
|
||||
|
||||
1. Add characterization tests around the existing external LLM prompt service if extraction is needed.
|
||||
2. Add the SEG crop image encoder and image-mode tests.
|
||||
3. Add response formatting logic and service tests.
|
||||
4. Add the application service with fake-boundary tests.
|
||||
5. Add the v3 node schema and execute forwarding tests.
|
||||
6. Register the v3 node and update registration tests.
|
||||
7. Add tooltips and update tooltip coverage tests.
|
||||
8. Run focused tests while developing.
|
||||
9. Run full verification gates before completion.
|
||||
|
||||
## Verification Commands
|
||||
|
||||
Run all commands from the repository root with PowerShell syntax.
|
||||
|
||||
Required gates:
|
||||
|
||||
```powershell
|
||||
..\..\venv\Scripts\ruff.exe format .
|
||||
..\..\venv\Scripts\ruff.exe check .
|
||||
..\..\venv\Scripts\mypy.exe --strict simple_syrup tests
|
||||
..\..\venv\Scripts\python.exe -m pytest -n auto -q
|
||||
```
|
||||
|
||||
Do not substitute global tools.
|
||||
|
||||
If a required tool is missing from `..\..\venv`, install or update development dependencies in that environment before verification. Do not create a local `.venv`.
|
||||
|
||||
## Definition Of Done
|
||||
|
||||
- The node is exported only through Comfy v3.
|
||||
- `SimpleSyrup.TagSEGSWithExternalLLM` appears in `get_nodes()`.
|
||||
- The node returns unchanged SEGS plus one positive conditioning entry per SEG.
|
||||
- One external LLM call is made per SEG.
|
||||
- `transparent mask`, `black mask`, and `full crop` are implemented and tested.
|
||||
- Formatting controls are implemented and tested.
|
||||
- All validation failures are explicit and actionable.
|
||||
- Tooltips cover every input and output.
|
||||
- Tests cover service behavior, image encoding, v3 schema, execution forwarding, registration, and tooltip coverage.
|
||||
- No legacy node class or legacy mapping export is introduced.
|
||||
- Required ruff, mypy, and pytest gates pass in `..\..\venv`.
|
||||
+3
-5
@@ -12,9 +12,8 @@ from . import simple_syrup as _simple_syrup_package
|
||||
|
||||
sys.modules.setdefault("simple_syrup", _simple_syrup_package)
|
||||
|
||||
from .simple_syrup.nodes import ( # noqa: E402
|
||||
NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS,
|
||||
from .simple_syrup.runtime.external_llm_routes import ( # noqa: E402
|
||||
register_external_llm_routes,
|
||||
)
|
||||
from .simple_syrup.runtime.settings_routes import register_settings_routes # noqa: E402
|
||||
|
||||
@@ -40,10 +39,9 @@ async def comfy_entrypoint() -> object:
|
||||
|
||||
|
||||
register_settings_routes()
|
||||
register_external_llm_routes()
|
||||
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"WEB_DIRECTORY",
|
||||
"comfy_entrypoint",
|
||||
]
|
||||
|
||||
Generated
+2
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.1.0",
|
||||
"version": "1.4.0",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.1.0",
|
||||
"version": "1.4.0",
|
||||
"license": "AGPL-3.0-or-later",
|
||||
"devDependencies": {
|
||||
"@eslint/js": "^9.39.1",
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "simple-syrup-comfyui",
|
||||
"version": "1.1.0",
|
||||
"version": "1.4.0",
|
||||
"private": true,
|
||||
"license": "AGPL-3.0-or-later",
|
||||
"type": "module",
|
||||
|
||||
+6
-1
@@ -5,7 +5,7 @@ build-backend = "setuptools.build_meta"
|
||||
[project]
|
||||
name = "SimpleSyrup"
|
||||
description = "Workflow-focused ComfyUI extensions for image generation."
|
||||
version = "1.1.0"
|
||||
version = "1.4.0"
|
||||
license = "AGPL-3.0-or-later"
|
||||
license-files = ["LICENSE"]
|
||||
requires-python = ">=3.11"
|
||||
@@ -65,3 +65,8 @@ ignore_missing_imports = true
|
||||
[tool.pytest.ini_options]
|
||||
pythonpath = [".", "../.."]
|
||||
testpaths = ["tests"]
|
||||
filterwarnings = [
|
||||
"error",
|
||||
"ignore:builtin type SwigPyPacked has no __module__ attribute:DeprecationWarning",
|
||||
"ignore:builtin type SwigPyObject has no __module__ attribute:DeprecationWarning",
|
||||
]
|
||||
|
||||
@@ -6,3 +6,4 @@ timm>=0.6.13
|
||||
addict>=2.4.0
|
||||
yapf>=0.43.0
|
||||
huggingface-hub>=0.34.0
|
||||
keyring>=25.0.0
|
||||
|
||||
@@ -6,6 +6,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
__version__ = "1.1.0"
|
||||
__version__ = "1.4.0"
|
||||
|
||||
__all__: list[str] = ["__version__"]
|
||||
|
||||
@@ -40,6 +40,23 @@ class ConditioningBatch:
|
||||
return ConditioningBatch((*self.entries, conditioning))
|
||||
|
||||
|
||||
def batch_conditioning(
|
||||
values: tuple[Conditioning | ConditioningBatch, ...],
|
||||
) -> ConditioningBatch:
|
||||
"""Flatten conditioning values and batches into one ordered batch."""
|
||||
|
||||
if not values:
|
||||
raise ValueError("Batch Region Conditioning requires one or more inputs.")
|
||||
|
||||
entries: list[Conditioning] = []
|
||||
for value in values:
|
||||
if isinstance(value, ConditioningBatch):
|
||||
entries.extend(value.entries)
|
||||
else:
|
||||
entries.append(value)
|
||||
return ConditioningBatch(tuple(entries))
|
||||
|
||||
|
||||
def split_prompt_batch(text: str, separator: str = "[SEP]") -> tuple[str, ...]:
|
||||
"""Split prompt text into ordered chunks using a configurable separator."""
|
||||
|
||||
|
||||
@@ -0,0 +1,153 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Domain objects and validation for external LLM provider integration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import ClassVar
|
||||
from urllib.parse import urlparse
|
||||
|
||||
DEFAULT_EXTERNAL_LLM_MAX_TOKENS = 1024
|
||||
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT = "default"
|
||||
EXTERNAL_LLM_REASONING_EFFORTS = (
|
||||
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
"high",
|
||||
"medium",
|
||||
"low",
|
||||
"off",
|
||||
)
|
||||
|
||||
|
||||
class ExternalLLMConfigError(ValueError):
|
||||
"""Raised when external LLM configuration is invalid."""
|
||||
|
||||
|
||||
class ExternalLLMProviderError(RuntimeError):
|
||||
"""Raised when an external LLM provider request fails."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExternalLLMConfig:
|
||||
"""Validated non-secret external LLM provider configuration."""
|
||||
|
||||
base_url: str
|
||||
cached_models: tuple[str, ...] = ()
|
||||
default_model: str = ""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate and normalize configuration values."""
|
||||
|
||||
normalized_url = normalize_base_url(self.base_url)
|
||||
models = normalize_model_ids(self.cached_models)
|
||||
default = self.default_model.strip()
|
||||
if default and models and default not in models:
|
||||
raise ExternalLLMConfigError(
|
||||
"Default external LLM model must be one of the cached models."
|
||||
)
|
||||
object.__setattr__(self, "base_url", normalized_url)
|
||||
object.__setattr__(self, "cached_models", models)
|
||||
object.__setattr__(self, "default_model", default)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExternalLLMModel:
|
||||
"""A provider model advertised through an OpenAI-compatible models response."""
|
||||
|
||||
id: str
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Reject empty provider model identifiers."""
|
||||
|
||||
model_id = self.id.strip()
|
||||
if not model_id:
|
||||
raise ExternalLLMConfigError("External LLM model id must not be empty.")
|
||||
object.__setattr__(self, "id", model_id)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExternalLLMChatRequest:
|
||||
"""Validated chat completion request fields for an external LLM provider."""
|
||||
|
||||
VALID_REASONING_EFFORTS: ClassVar[tuple[str, ...]] = EXTERNAL_LLM_REASONING_EFFORTS
|
||||
|
||||
model: str
|
||||
system_prompt: str
|
||||
user_prompt: str
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT
|
||||
image_data_url: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
"""Validate prompt request values before provider execution."""
|
||||
|
||||
model = self.model.strip()
|
||||
reasoning_effort = self.reasoning_effort.strip()
|
||||
if not model:
|
||||
raise ExternalLLMConfigError("External LLM model must not be empty.")
|
||||
if not self.user_prompt.strip():
|
||||
raise ExternalLLMConfigError(
|
||||
"User prompt must not be empty for external LLM requests."
|
||||
)
|
||||
if self.max_tokens < 1:
|
||||
raise ExternalLLMConfigError("External LLM max tokens must be at least 1.")
|
||||
if reasoning_effort not in self.VALID_REASONING_EFFORTS:
|
||||
choices = ", ".join(self.VALID_REASONING_EFFORTS)
|
||||
raise ExternalLLMConfigError(
|
||||
f"External LLM reasoning effort must be one of: {choices}."
|
||||
)
|
||||
if self.image_data_url is not None and not self.image_data_url.startswith(
|
||||
"data:image/"
|
||||
):
|
||||
raise ExternalLLMConfigError(
|
||||
"External LLM image input must be an image data URL."
|
||||
)
|
||||
object.__setattr__(self, "model", model)
|
||||
object.__setattr__(self, "reasoning_effort", reasoning_effort)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExternalLLMChatResponse:
|
||||
"""Assistant response content returned by an external LLM provider."""
|
||||
|
||||
content: str
|
||||
|
||||
|
||||
def normalize_base_url(base_url: str) -> str:
|
||||
"""Return a normalized absolute HTTP(S) endpoint URL."""
|
||||
|
||||
candidate = base_url.strip().rstrip("/")
|
||||
if not candidate:
|
||||
raise ExternalLLMConfigError(
|
||||
"Configure an external LLM endpoint in SimpleSyrup settings before "
|
||||
"using this node."
|
||||
)
|
||||
|
||||
parsed = urlparse(candidate)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
||||
raise ExternalLLMConfigError(
|
||||
"External LLM endpoint must be an absolute http:// or https:// URL."
|
||||
)
|
||||
return candidate
|
||||
|
||||
|
||||
def normalize_model_ids(values: object) -> tuple[str, ...]:
|
||||
"""Return non-empty de-duplicated model ids in provider order."""
|
||||
|
||||
if not isinstance(values, (list, tuple)):
|
||||
raise ExternalLLMConfigError("External LLM cached models must be a list.")
|
||||
|
||||
models: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for value in values:
|
||||
if not isinstance(value, str):
|
||||
raise ExternalLLMConfigError(
|
||||
"External LLM cached model ids must be strings."
|
||||
)
|
||||
model = value.strip()
|
||||
if model and model not in seen:
|
||||
seen.add(model)
|
||||
models.append(model)
|
||||
return tuple(models)
|
||||
@@ -0,0 +1,76 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Prepare Prompt-Control prompt text for scheduling and encoding."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .conditioning_batch import split_prompt_batch
|
||||
|
||||
PROMPT_TEXT_PATTERN = r"(?:^|>)([^<]+)(?=<|$)"
|
||||
LORA_TAG_PATTERN = r"<[^>]*>"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PreparedPromptChunk:
|
||||
"""Store one prompt chunk's cleaned text and scheduling tags."""
|
||||
|
||||
text: str
|
||||
lora_tags: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PreparedPromptSide:
|
||||
"""Store ordered prompt chunks and all scheduling tags for one prompt side."""
|
||||
|
||||
chunks: tuple[PreparedPromptChunk, ...]
|
||||
lora_tags: str
|
||||
|
||||
|
||||
def extract_prompt_text(text: str) -> str:
|
||||
"""Return prompt text outside angle-bracket Prompt-Control tags."""
|
||||
|
||||
return _extract_all_matches(text, PROMPT_TEXT_PATTERN)
|
||||
|
||||
|
||||
def extract_lora_tags(text: str) -> str:
|
||||
"""Return angle-bracket Prompt-Control tags joined with newlines."""
|
||||
|
||||
return _extract_all_matches(text, LORA_TAG_PATTERN)
|
||||
|
||||
|
||||
def prepare_prompt_side(text: str, separator: str) -> PreparedPromptSide:
|
||||
"""Split a prompt side into cleaned chunks and aggregate LoRA tags."""
|
||||
|
||||
chunks = tuple(
|
||||
PreparedPromptChunk(
|
||||
text=extract_prompt_text(chunk),
|
||||
lora_tags=extract_lora_tags(chunk),
|
||||
)
|
||||
for chunk in split_prompt_batch(text, separator)
|
||||
)
|
||||
lora_tags = "\n".join(chunk.lora_tags for chunk in chunks if chunk.lora_tags)
|
||||
return PreparedPromptSide(chunks=chunks, lora_tags=lora_tags)
|
||||
|
||||
|
||||
def apply_encode_style(encode_style: str, prompt_text: str) -> str:
|
||||
"""Prepend Prompt-Control encode style text exactly as provided."""
|
||||
|
||||
if not encode_style:
|
||||
return prompt_text
|
||||
return f"{encode_style}{prompt_text}"
|
||||
|
||||
|
||||
def _extract_all_matches(text: str, pattern: str) -> str:
|
||||
"""Match Comfy's RegexExtract All Matches behavior."""
|
||||
|
||||
matches = re.findall(pattern, text, re.IGNORECASE)
|
||||
if not matches:
|
||||
return ""
|
||||
if isinstance(matches[0], tuple):
|
||||
return "\n".join(match[0] for match in matches)
|
||||
return "\n".join(matches)
|
||||
@@ -184,6 +184,32 @@ def to_impact_compatible_segs_group(segs_group: NativeSegsGroup) -> list[ImpactS
|
||||
return [to_impact_compatible_segs(segs) for segs in segs_group]
|
||||
|
||||
|
||||
def batch_segs(values: Iterable[object]) -> NativeSegs:
|
||||
"""Return one SEGS payload containing all segments in input order."""
|
||||
|
||||
raw_values = tuple(values)
|
||||
if not raw_values:
|
||||
raise ValueError("Batch SEGS requires one or more SEGS inputs.")
|
||||
|
||||
expected_header: SegsHeader | None = None
|
||||
batched_segments: list[Segment] = []
|
||||
for index, value in enumerate(raw_values, start=1):
|
||||
header, segments = coerce_segs(value)
|
||||
if expected_header is None:
|
||||
expected_header = header
|
||||
elif header != expected_header:
|
||||
raise ValueError(
|
||||
"Batch SEGS requires all SEGS inputs to use the same image size; "
|
||||
f"input {index} is {_format_header(header)} but input 1 is "
|
||||
f"{_format_header(expected_header)}."
|
||||
)
|
||||
batched_segments.extend(segments)
|
||||
|
||||
if expected_header is None:
|
||||
raise ValueError("Batch SEGS requires one or more SEGS inputs.")
|
||||
return expected_header, tuple(batched_segments)
|
||||
|
||||
|
||||
def limit_segs(segs: NativeSegs, keep_only: int, keep_by: str) -> NativeSegs:
|
||||
"""Return SEGS limited by a user-facing ranking policy."""
|
||||
|
||||
@@ -284,6 +310,13 @@ def _coerce_header(value: object) -> SegsHeader:
|
||||
return height, width
|
||||
|
||||
|
||||
def _format_header(header: SegsHeader) -> str:
|
||||
"""Return a height-first image size description."""
|
||||
|
||||
height, width = header
|
||||
return f"{height}x{width}"
|
||||
|
||||
|
||||
def _looks_like_segs(value: object) -> bool:
|
||||
"""Return whether a value has the outer shape of one SEGS payload."""
|
||||
|
||||
|
||||
@@ -73,7 +73,8 @@ def build_tiled_diffusion_plan(
|
||||
)
|
||||
effective_tile_width = min(tile_width, latent_width)
|
||||
effective_tile_height = min(tile_height, latent_height)
|
||||
effective_overlap = max(0, min(overlap, min(tile_width, tile_height) - 4))
|
||||
max_effective_overlap = min(effective_tile_width, effective_tile_height) - 4
|
||||
effective_overlap = max(0, min(overlap, max_effective_overlap))
|
||||
|
||||
tiles = _split_tiles(
|
||||
latent_width=latent_width,
|
||||
|
||||
@@ -2,136 +2,7 @@
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI node registration for SimpleSyrup."""
|
||||
"""Implementation modules for SimpleSyrup node behavior.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from .conditioning_batch_pack import ConditioningBatchAppend, ConditioningBatchStart
|
||||
from .detail_segs_as_regions import DetailSEGSAsRegions
|
||||
from .detail_segs_by_scale_factor import DetailSEGSByScaleFactor
|
||||
from .detail_segs_by_scale_factor_tiled_diffusion import (
|
||||
DetailSEGSByScaleFactorTiledDiffusion,
|
||||
)
|
||||
from .detect_segs_with_ultralytics import DetectSEGSWithUltralytics
|
||||
from .encode_prompt_batch import EncodePromptBatch
|
||||
from .grounded_sam_model_info import GroundedSAMModelInfo
|
||||
from .grounding_dino_model_loader import GroundingDINOModelLoader
|
||||
from .image_resize_to_target import ResizeImageToTarget
|
||||
from .ksampler_extras import KSamplerExtras
|
||||
from .ksampler_tiled_diffusion import KSamplerTiledDiffusion
|
||||
from .latent_diagnostics import LatentDiagnostics
|
||||
from .layerstyle_sam_models_adapter import LayerStyleSAMModelsAdapter
|
||||
from .load_ultralytics_model import LoadUltralyticsModel
|
||||
from .prompt_encode_style import PromptEncodeStyle
|
||||
from .prompt_encode_style_and_normalization import PromptEncodeStyleAndNormalization
|
||||
from .prompt_segs_with_sam import PromptSEGSWithSAM
|
||||
from .provenance_latent import SimpleVAEEncode, UpscaleLatentFromImage
|
||||
from .sam_model_loader import SAMModelLoader
|
||||
from .scale_factor import ScaleFactor
|
||||
from .seed import Seed
|
||||
from .simple_load_anima import SimpleLoadAnima
|
||||
from .simple_load_checkpoint import SimpleLoadCheckpoint
|
||||
from .tile_and_tag_segs import TileAndTagSEGS
|
||||
from .vitmatte_model_loader import ViTMatteModelLoader
|
||||
from .wd14_tagger_loader import WD14TaggerLoader
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SimpleSyrup.ConditioningBatchAppend": ConditioningBatchAppend,
|
||||
"SimpleSyrup.ConditioningBatchStart": ConditioningBatchStart,
|
||||
"SimpleSyrup.GroundedSAMModelInfo": GroundedSAMModelInfo,
|
||||
"SimpleSyrup.GroundingDINOModelLoader": GroundingDINOModelLoader,
|
||||
"SimpleSyrup.KSamplerExtras": KSamplerExtras,
|
||||
"SimpleSyrup.KSamplerTiledDiffusion": KSamplerTiledDiffusion,
|
||||
"SimpleSyrup.LayerStyleSAMModelsAdapter": LayerStyleSAMModelsAdapter,
|
||||
"SimpleSyrup.LatentDiagnostics": LatentDiagnostics,
|
||||
"SimpleSyrup.PromptEncodeStyle": PromptEncodeStyle,
|
||||
"SimpleSyrup.PromptEncodeStyleAndNormalization": PromptEncodeStyleAndNormalization,
|
||||
"SimpleSyrup.PromptSEGSWithSAM": PromptSEGSWithSAM,
|
||||
"SimpleSyrup.SimpleVAEEncode": SimpleVAEEncode,
|
||||
"SimpleSyrup.UpscaleLatentFromImage": UpscaleLatentFromImage,
|
||||
"SimpleSyrup.ResizeImageToTarget": ResizeImageToTarget,
|
||||
"SimpleSyrup.DetailSEGSAsRegions": DetailSEGSAsRegions,
|
||||
"SimpleSyrup.DetailSEGSByScaleFactor": DetailSEGSByScaleFactor,
|
||||
"SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion": (
|
||||
DetailSEGSByScaleFactorTiledDiffusion
|
||||
),
|
||||
"SimpleSyrup.SAMModelLoader": SAMModelLoader,
|
||||
"SimpleSyrup.ScaleFactor": ScaleFactor,
|
||||
"SimpleSyrup.Seed": Seed,
|
||||
"SimpleSyrup.SimpleLoadAnima": SimpleLoadAnima,
|
||||
"SimpleSyrup.SimpleLoadCheckpoint": SimpleLoadCheckpoint,
|
||||
"SimpleSyrup.LoadUltralyticsModel": LoadUltralyticsModel,
|
||||
"SimpleSyrup.DetectSEGSWithUltralytics": DetectSEGSWithUltralytics,
|
||||
"SimpleSyrup.EncodePromptBatch": EncodePromptBatch,
|
||||
"SimpleSyrup.TileAndTagSEGS": TileAndTagSEGS,
|
||||
"SimpleSyrup.ViTMatteModelLoader": ViTMatteModelLoader,
|
||||
"SimpleSyrup.WD14TaggerLoader": WD14TaggerLoader,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"SimpleSyrup.ConditioningBatchAppend": "Conditioning Batch Append",
|
||||
"SimpleSyrup.ConditioningBatchStart": "Conditioning Batch Start",
|
||||
"SimpleSyrup.GroundedSAMModelInfo": "Grounded SAM Model Info",
|
||||
"SimpleSyrup.GroundingDINOModelLoader": "GroundingDINO Model Loader",
|
||||
"SimpleSyrup.KSamplerExtras": "KSampler (Extras)",
|
||||
"SimpleSyrup.KSamplerTiledDiffusion": "KSampler (Tiled Diffusion)",
|
||||
"SimpleSyrup.LayerStyleSAMModelsAdapter": "LayerStyle SAM Models Adapter",
|
||||
"SimpleSyrup.LatentDiagnostics": "Latent Diagnostics",
|
||||
"SimpleSyrup.PromptEncodeStyle": "Prompt Encode Style",
|
||||
"SimpleSyrup.PromptEncodeStyleAndNormalization": (
|
||||
"Prompt Encode Style & Normalization"
|
||||
),
|
||||
"SimpleSyrup.PromptSEGSWithSAM": "Prompt SEGS w/ SAM",
|
||||
"SimpleSyrup.SimpleVAEEncode": "Simple VAE Encode",
|
||||
"SimpleSyrup.UpscaleLatentFromImage": "Upscale Latent From Image",
|
||||
"SimpleSyrup.ResizeImageToTarget": "Resize Image to Target",
|
||||
"SimpleSyrup.DetailSEGSAsRegions": "Detail SEGS as Regions",
|
||||
"SimpleSyrup.DetailSEGSByScaleFactor": "Detail SEGS by Scale Factor",
|
||||
"SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion": (
|
||||
"Detail SEGS by Scale Factor w/ Tiled Diffusion"
|
||||
),
|
||||
"SimpleSyrup.SAMModelLoader": "SAM Model Loader",
|
||||
"SimpleSyrup.ScaleFactor": "Scale Factor",
|
||||
"SimpleSyrup.Seed": "Seed",
|
||||
"SimpleSyrup.SimpleLoadAnima": "Simple Load Anima",
|
||||
"SimpleSyrup.SimpleLoadCheckpoint": "Simple Load Checkpoint",
|
||||
"SimpleSyrup.LoadUltralyticsModel": "Load Ultralytics Model",
|
||||
"SimpleSyrup.DetectSEGSWithUltralytics": "Detect SEGS w/ Ultralytics",
|
||||
"SimpleSyrup.EncodePromptBatch": "Encode Prompt Batch",
|
||||
"SimpleSyrup.TileAndTagSEGS": "Tile & Tag SEGS",
|
||||
"SimpleSyrup.ViTMatteModelLoader": "ViTMatte Model Loader",
|
||||
"SimpleSyrup.WD14TaggerLoader": "Load WD14 Tagger",
|
||||
}
|
||||
|
||||
__all__ = [
|
||||
"ConditioningBatchAppend",
|
||||
"ConditioningBatchStart",
|
||||
"GroundedSAMModelInfo",
|
||||
"GroundingDINOModelLoader",
|
||||
"KSamplerExtras",
|
||||
"KSamplerTiledDiffusion",
|
||||
"LayerStyleSAMModelsAdapter",
|
||||
"LatentDiagnostics",
|
||||
"DetectSEGSWithUltralytics",
|
||||
"DetailSEGSAsRegions",
|
||||
"DetailSEGSByScaleFactor",
|
||||
"DetailSEGSByScaleFactorTiledDiffusion",
|
||||
"EncodePromptBatch",
|
||||
"LoadUltralyticsModel",
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_DISPLAY_NAME_MAPPINGS",
|
||||
"PromptEncodeStyle",
|
||||
"PromptEncodeStyleAndNormalization",
|
||||
"PromptSEGSWithSAM",
|
||||
"ResizeImageToTarget",
|
||||
"SAMModelLoader",
|
||||
"ScaleFactor",
|
||||
"Seed",
|
||||
"SimpleLoadAnima",
|
||||
"SimpleLoadCheckpoint",
|
||||
"SimpleVAEEncode",
|
||||
"TileAndTagSEGS",
|
||||
"UpscaleLatentFromImage",
|
||||
"ViTMatteModelLoader",
|
||||
"WD14TaggerLoader",
|
||||
]
|
||||
ComfyUI registration is v3-only and lives in `simple_syrup.nodes_v3`.
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI node declaration for batching regional conditioning."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..domain.conditioning_batch import batch_conditioning
|
||||
from ..nodes import tooltips
|
||||
|
||||
|
||||
class BatchRegionConditioning:
|
||||
"""Combine conditioning and conditioning batches for regional detailing."""
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING_BATCH",)
|
||||
RETURN_NAMES = ("batch",)
|
||||
OUTPUT_TOOLTIPS = (tooltips.BATCH_REGION_CONDITIONING_OUTPUT,)
|
||||
FUNCTION = "batch"
|
||||
CATEGORY = "SimpleSyrup/Conditioning"
|
||||
DESCRIPTION = (
|
||||
"Combines two CONDITIONING or CONDITIONING_BATCH inputs into one ordered "
|
||||
"regional conditioning batch. Chain this node to batch more sources."
|
||||
)
|
||||
SEARCH_ALIASES = [
|
||||
"batch",
|
||||
"conditioning batch",
|
||||
"region conditioning",
|
||||
"segs prompts",
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare legacy ComfyUI inputs for regional conditioning batching."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"first": (
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.BATCH_REGION_CONDITIONING_FIRST},
|
||||
),
|
||||
"second": (
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.BATCH_REGION_CONDITIONING_SECOND},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def batch(self, first: Any, second: Any) -> tuple[object]:
|
||||
"""Batch two conditioning inputs in input order."""
|
||||
|
||||
return (batch_conditioning((first, second)),)
|
||||
@@ -0,0 +1,44 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI node declaration for batching SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..domain.segs import batch_segs, to_impact_compatible_segs
|
||||
from ..nodes import tooltips
|
||||
|
||||
|
||||
class BatchSEGS:
|
||||
"""Combine two SEGS payloads into one ordered SEGS payload."""
|
||||
|
||||
RETURN_TYPES = ("SEGS",)
|
||||
RETURN_NAMES = ("segs",)
|
||||
OUTPUT_TOOLTIPS = (tooltips.BATCH_SEGS_OUTPUT,)
|
||||
FUNCTION = "batch"
|
||||
CATEGORY = "SimpleSyrup/Detection"
|
||||
DESCRIPTION = (
|
||||
"Combines two SEGS inputs into one ordered SEGS payload. Chain this node "
|
||||
"to batch more than two SEGS sources."
|
||||
)
|
||||
SEARCH_ALIASES = ["batch", "merge", "join", "combine", "segs"]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare legacy ComfyUI inputs for SEGS batching."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"first": ("SEGS", {"tooltip": tooltips.BATCH_SEGS_FIRST}),
|
||||
"second": ("SEGS", {"tooltip": tooltips.BATCH_SEGS_SECOND}),
|
||||
},
|
||||
}
|
||||
|
||||
def batch(self, first: object, second: object) -> tuple[object]:
|
||||
"""Batch two SEGS payloads in input order."""
|
||||
|
||||
native = batch_segs((first, second))
|
||||
return (to_impact_compatible_segs(native),)
|
||||
@@ -0,0 +1,111 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI node declaration for external LLM prompt generation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..domain.external_llm import (
|
||||
DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
EXTERNAL_LLM_REASONING_EFFORTS,
|
||||
)
|
||||
from ..services.external_llm_prompt_service import ExternalLLMPromptService
|
||||
from . import tooltips
|
||||
|
||||
MAX_EXTERNAL_LLM_MAX_TOKENS = 32768
|
||||
|
||||
|
||||
class ExternalLLMPrompt:
|
||||
"""Expose an OpenAI-compatible external LLM prompt request as a string node."""
|
||||
|
||||
_service = ExternalLLMPromptService()
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("response",)
|
||||
OUTPUT_TOOLTIPS = (tooltips.EXTERNAL_LLM_RESPONSE_OUTPUT,)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "SimpleSyrup/Prompting"
|
||||
DESCRIPTION = "Sends system and user prompts to a configured external LLM provider."
|
||||
SEARCH_ALIASES = ["llm", "openai", "prompt", "preprocess", "text"]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare cached external LLM prompt inputs."""
|
||||
|
||||
choices = cls._service.model_choices()
|
||||
return {
|
||||
"required": {
|
||||
"model": (
|
||||
choices,
|
||||
{
|
||||
"default": choices[0],
|
||||
"tooltip": tooltips.EXTERNAL_LLM_MODEL_INPUT,
|
||||
},
|
||||
),
|
||||
"system_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": tooltips.EXTERNAL_LLM_SYSTEM_PROMPT_INPUT,
|
||||
},
|
||||
),
|
||||
"user_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": tooltips.EXTERNAL_LLM_USER_PROMPT_INPUT,
|
||||
},
|
||||
),
|
||||
"max_tokens": (
|
||||
"INT",
|
||||
{
|
||||
"default": DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
"min": 1,
|
||||
"max": MAX_EXTERNAL_LLM_MAX_TOKENS,
|
||||
"step": 1,
|
||||
"tooltip": tooltips.EXTERNAL_LLM_MAX_TOKENS_INPUT,
|
||||
},
|
||||
),
|
||||
"reasoning_effort": (
|
||||
list(EXTERNAL_LLM_REASONING_EFFORTS),
|
||||
{
|
||||
"default": DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
"tooltip": tooltips.EXTERNAL_LLM_REASONING_EFFORT_INPUT,
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"image": (
|
||||
"IMAGE",
|
||||
{"tooltip": tooltips.EXTERNAL_LLM_IMAGE_INPUT},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def generate(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
image: object | None = None,
|
||||
) -> tuple[str]:
|
||||
"""Return the external LLM assistant response."""
|
||||
|
||||
return (
|
||||
self._service.generate(
|
||||
model,
|
||||
system_prompt,
|
||||
user_prompt,
|
||||
max_tokens,
|
||||
reasoning_effort,
|
||||
image=image,
|
||||
),
|
||||
)
|
||||
@@ -9,6 +9,9 @@ from __future__ import annotations
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.conditioning_batch import ConditioningBatch, select_conditioning
|
||||
from ..runtime import sampling_samplers, sampling_schedulers
|
||||
from . import tooltips
|
||||
|
||||
@@ -74,11 +77,11 @@ class KSamplerExtras:
|
||||
{"tooltip": tooltips.SCHEDULER},
|
||||
),
|
||||
"positive": (
|
||||
"CONDITIONING",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.POSITIVE_CONDITIONING},
|
||||
),
|
||||
"negative": (
|
||||
"CONDITIONING",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.NEGATIVE_CONDITIONING},
|
||||
),
|
||||
"latent_image": ("LATENT", {"tooltip": tooltips.LATENT_IMAGE}),
|
||||
@@ -138,20 +141,37 @@ class KSamplerExtras:
|
||||
|
||||
callback = latent_preview.prepare_callback(model, steps)
|
||||
disable_pbar = not comfy_utils.PROGRESS_BAR_ENABLED
|
||||
samples = comfy_sample.sample_custom(
|
||||
model,
|
||||
noise,
|
||||
cfg,
|
||||
sampler,
|
||||
sigmas,
|
||||
positive,
|
||||
negative,
|
||||
latent_samples,
|
||||
noise_mask=noise_mask,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=seed,
|
||||
)
|
||||
if _uses_conditioning_batch(positive, negative):
|
||||
samples = _sample_conditioning_batch(
|
||||
comfy_sample=comfy_sample,
|
||||
model=model,
|
||||
noise=noise,
|
||||
cfg=cfg,
|
||||
sampler=sampler,
|
||||
sigmas=sigmas,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_samples=latent_samples,
|
||||
noise_mask=noise_mask,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=seed,
|
||||
)
|
||||
else:
|
||||
samples = comfy_sample.sample_custom(
|
||||
model,
|
||||
noise,
|
||||
cfg,
|
||||
sampler,
|
||||
sigmas,
|
||||
positive,
|
||||
negative,
|
||||
latent_samples,
|
||||
noise_mask=noise_mask,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=seed,
|
||||
)
|
||||
|
||||
output = latent_image.copy()
|
||||
output.pop("downscale_ratio_spacial", None)
|
||||
@@ -167,6 +187,68 @@ def _comfy_sample() -> Any:
|
||||
return comfy.sample
|
||||
|
||||
|
||||
def _uses_conditioning_batch(positive: Any, negative: Any) -> bool:
|
||||
"""Return whether either conditioning input needs per-item selection."""
|
||||
|
||||
return isinstance(positive, ConditioningBatch) or isinstance(
|
||||
negative,
|
||||
ConditioningBatch,
|
||||
)
|
||||
|
||||
|
||||
def _sample_conditioning_batch(
|
||||
*,
|
||||
comfy_sample: Any,
|
||||
model: Any,
|
||||
noise: torch.Tensor,
|
||||
cfg: float,
|
||||
sampler: Any,
|
||||
sigmas: torch.Tensor,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
latent_samples: torch.Tensor,
|
||||
noise_mask: Any,
|
||||
callback: Any,
|
||||
disable_pbar: bool,
|
||||
seed: int,
|
||||
) -> torch.Tensor:
|
||||
"""Sample each latent batch item with its selected conditioning."""
|
||||
|
||||
sampled: list[torch.Tensor] = []
|
||||
for index in range(int(latent_samples.shape[0])):
|
||||
sampled.append(
|
||||
comfy_sample.sample_custom(
|
||||
model,
|
||||
noise[index : index + 1],
|
||||
cfg,
|
||||
sampler,
|
||||
sigmas,
|
||||
select_conditioning(positive, index),
|
||||
select_conditioning(negative, index),
|
||||
latent_samples[index : index + 1],
|
||||
noise_mask=_slice_noise_mask(noise_mask, index, latent_samples),
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=seed,
|
||||
)
|
||||
)
|
||||
return torch.cat(sampled, dim=0)
|
||||
|
||||
|
||||
def _slice_noise_mask(
|
||||
noise_mask: Any,
|
||||
index: int,
|
||||
latent_samples: torch.Tensor,
|
||||
) -> Any:
|
||||
"""Return the noise mask slice matching one latent batch item."""
|
||||
|
||||
if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int(
|
||||
latent_samples.shape[0],
|
||||
):
|
||||
return noise_mask[index : index + 1]
|
||||
return noise_mask
|
||||
|
||||
|
||||
def _comfy_utils() -> Any:
|
||||
"""Import ComfyUI utility state lazily."""
|
||||
|
||||
|
||||
@@ -84,11 +84,11 @@ class KSamplerTiledDiffusion:
|
||||
{"tooltip": tooltips.SCHEDULER},
|
||||
),
|
||||
"positive": (
|
||||
"CONDITIONING",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.POSITIVE_CONDITIONING},
|
||||
),
|
||||
"negative": (
|
||||
"CONDITIONING",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
{"tooltip": tooltips.NEGATIVE_CONDITIONING},
|
||||
),
|
||||
"latent_image": ("LATENT", {"tooltip": tooltips.LATENT_IMAGE}),
|
||||
|
||||
@@ -15,8 +15,8 @@ class PromptEncodeStyle:
|
||||
"""Build Prompt Control STYLE tags from encode-style selections."""
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("style_tag",)
|
||||
OUTPUT_TOOLTIPS = ("Prompt Control STYLE tag text for prompt encoding workflows.",)
|
||||
RETURN_NAMES = ("encode_style",)
|
||||
OUTPUT_TOOLTIPS = ("Prompt Control encode style text for prompt workflows.",)
|
||||
FUNCTION = "build"
|
||||
CATEGORY = "SimpleSyrup/Prompt"
|
||||
DESCRIPTION = "Builds a Prompt Control STYLE tag from an encode-style selection."
|
||||
|
||||
@@ -19,9 +19,9 @@ class PromptEncodeStyleAndNormalization:
|
||||
"""Build STYLE tags from encode-style and normalization selections."""
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("style_tag",)
|
||||
RETURN_NAMES = ("encode_style",)
|
||||
OUTPUT_TOOLTIPS = (
|
||||
"Prompt Control STYLE tag text with the selected normalization behavior.",
|
||||
"Prompt Control encode style text with the selected normalization behavior.",
|
||||
)
|
||||
FUNCTION = "build"
|
||||
CATEGORY = "SimpleSyrup/Prompt"
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Legacy ComfyUI node for Prompt-Control prompt scheduling and encoding."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..runtime.prompt_control_schedule_encode_graph import (
|
||||
PromptControlScheduleEncodeGraphBuilder,
|
||||
)
|
||||
|
||||
|
||||
class ScheduleAndEncodePromptsWithPromptControl:
|
||||
"""Schedule Prompt-Control LoRAs and encode prompts with optional batches."""
|
||||
|
||||
RETURN_TYPES = (
|
||||
"MODEL",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
)
|
||||
RETURN_NAMES = ("model", "positive", "negative")
|
||||
OUTPUT_TOOLTIPS = (
|
||||
"Model after LoRA tags from positive and negative prompts are scheduled.",
|
||||
"Positive conditioning or SimpleSyrup conditioning batch.",
|
||||
"Negative conditioning or SimpleSyrup conditioning batch.",
|
||||
)
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = "SimpleSyrup/Conditioning"
|
||||
DESCRIPTION = (
|
||||
"Schedules Prompt-Control LoRAs and encodes prompts, using [SEP] to "
|
||||
"create SimpleSyrup conditioning batches."
|
||||
)
|
||||
SEARCH_ALIASES = ["prompt control", "schedule prompts", "encode prompts"]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare Prompt-Control schedule and encode inputs."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model": (
|
||||
"MODEL",
|
||||
{
|
||||
"rawLink": True,
|
||||
"tooltip": (
|
||||
"Model that receives LoRA changes found in "
|
||||
"Prompt-Control prompt tags."
|
||||
),
|
||||
},
|
||||
),
|
||||
"clip": (
|
||||
"CLIP",
|
||||
{
|
||||
"rawLink": True,
|
||||
"tooltip": (
|
||||
"CLIP connection used for scheduled hooks and "
|
||||
"cleaned prompt encoding."
|
||||
),
|
||||
},
|
||||
),
|
||||
"positive_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": (
|
||||
"Positive Prompt-Control text. [SEP] creates a "
|
||||
"conditioning batch for SimpleSyrup batch-aware nodes."
|
||||
),
|
||||
},
|
||||
),
|
||||
"negative_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": (
|
||||
"Negative Prompt-Control text. [SEP] creates a "
|
||||
"conditioning batch for SimpleSyrup batch-aware nodes."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"encode_style": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"forceInput": True,
|
||||
"tooltip": (
|
||||
"Encode style text from Prompt Encode Style or "
|
||||
"Prompt Encode Style & Normalization."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def execute(
|
||||
self,
|
||||
model: Any,
|
||||
clip: Any,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
encode_style: str = "",
|
||||
) -> Any:
|
||||
"""Build lazy Prompt-Control graph expansion for prompts."""
|
||||
|
||||
return PromptControlScheduleEncodeGraphBuilder().build(
|
||||
model=model,
|
||||
clip=clip,
|
||||
positive_prompt=positive_prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
encode_style=encode_style,
|
||||
)
|
||||
@@ -0,0 +1,172 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI node declaration for tagging existing SEGS with WD14."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..runtime.wd14_tagger import WD14TagFormattingControls
|
||||
from ..services.tag_segs_with_wd14_service import TagSEGSWithWD14Service
|
||||
from .tile_and_tag_segs import DEFAULT_EXCLUDE_TAGS
|
||||
|
||||
|
||||
class TagSEGSWithWD14:
|
||||
"""Create WD14 conditioning for existing SEGS."""
|
||||
|
||||
RETURN_TYPES = ("SEGS", "CONDITIONING_BATCH")
|
||||
RETURN_NAMES = ("segs", "positive")
|
||||
OUTPUT_TOOLTIPS = (
|
||||
tooltips.TAG_SEGS_SEGS_OUTPUT,
|
||||
tooltips.TAG_SEGS_POSITIVE_OUTPUT,
|
||||
)
|
||||
FUNCTION = "tag"
|
||||
CATEGORY = "SimpleSyrup/Detailing"
|
||||
DESCRIPTION = (
|
||||
"Tags existing SEGS crops with a connected WD14 tagger and returns "
|
||||
"aligned conditioning for SEGS detailing."
|
||||
)
|
||||
SEARCH_ALIASES = ["tag", "wd14", "segs", "detail", "regional"]
|
||||
|
||||
service_class: ClassVar[Callable[[], TagSEGSWithWD14Service]] = (
|
||||
TagSEGSWithWD14Service
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare ComfyUI inputs for WD14 tagging of existing SEGS."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", {"tooltip": tooltips.TAG_SEGS_IMAGE}),
|
||||
"segs": ("SEGS", {"tooltip": tooltips.TAG_SEGS_SEGS}),
|
||||
"clip": ("CLIP", {"tooltip": tooltips.TAG_SEGS_CLIP}),
|
||||
"wd14_tagger": (
|
||||
"WD14_TAGGER",
|
||||
{"tooltip": tooltips.TAG_SEGS_WD14_TAGGER},
|
||||
),
|
||||
"universal_positive": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": False,
|
||||
"tooltip": tooltips.TAG_SEGS_UNIVERSAL_POSITIVE,
|
||||
},
|
||||
),
|
||||
"threshold": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.35,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": tooltips.TILE_THRESHOLD,
|
||||
},
|
||||
),
|
||||
"character_threshold": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.05,
|
||||
"tooltip": tooltips.TILE_CHARACTER_THRESHOLD,
|
||||
},
|
||||
),
|
||||
"replace_underscore": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": True,
|
||||
"tooltip": tooltips.TILE_REPLACE_UNDERSCORE,
|
||||
},
|
||||
),
|
||||
"trailing_comma": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": tooltips.TILE_TRAILING_COMMA,
|
||||
},
|
||||
),
|
||||
"exclude_tags": (
|
||||
"STRING",
|
||||
{
|
||||
"default": DEFAULT_EXCLUDE_TAGS,
|
||||
"multiline": False,
|
||||
"tooltip": tooltips.TILE_EXCLUDE_TAGS,
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
def tag(
|
||||
self,
|
||||
image: object,
|
||||
segs: object,
|
||||
clip: Any,
|
||||
wd14_tagger: object,
|
||||
universal_positive: object,
|
||||
threshold: object,
|
||||
character_threshold: object,
|
||||
replace_underscore: object,
|
||||
trailing_comma: object,
|
||||
exclude_tags: object,
|
||||
) -> tuple[object, object]:
|
||||
"""Tag existing SEGS and return aligned conditioning."""
|
||||
|
||||
tag_controls = WD14TagFormattingControls(
|
||||
threshold=_float_input(threshold, "threshold"),
|
||||
character_threshold=_float_input(
|
||||
character_threshold,
|
||||
"character_threshold",
|
||||
),
|
||||
replace_underscore=_bool_input(
|
||||
replace_underscore,
|
||||
"replace_underscore",
|
||||
),
|
||||
trailing_comma=_bool_input(trailing_comma, "trailing_comma"),
|
||||
exclude_tags=_str_input(exclude_tags, "exclude_tags"),
|
||||
)
|
||||
result = (
|
||||
type(self)
|
||||
.service_class()
|
||||
.tag(
|
||||
image=image,
|
||||
segs=segs,
|
||||
clip=clip,
|
||||
wd14_tagger=wd14_tagger,
|
||||
tag_controls=tag_controls,
|
||||
universal_positive=_str_input(
|
||||
universal_positive,
|
||||
"universal_positive",
|
||||
),
|
||||
)
|
||||
)
|
||||
return result.segs, result.positive
|
||||
|
||||
|
||||
def _float_input(value: object, name: str) -> float:
|
||||
"""Return a float node input."""
|
||||
|
||||
if isinstance(value, (int, float, str)):
|
||||
return float(value)
|
||||
raise TypeError(f"Tag SEGS w/ WD14 requires '{name}' to be a float.")
|
||||
|
||||
|
||||
def _str_input(value: object, name: str) -> str:
|
||||
"""Return a string node input."""
|
||||
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
raise TypeError(f"Tag SEGS w/ WD14 requires '{name}' to be a string.")
|
||||
|
||||
|
||||
def _bool_input(value: object, name: str) -> bool:
|
||||
"""Return a boolean node input."""
|
||||
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
raise TypeError(f"Tag SEGS w/ WD14 requires '{name}' to be a boolean.")
|
||||
@@ -23,6 +23,27 @@ MODEL_OUTPUT = "Loaded diffusion model for downstream MODEL inputs."
|
||||
CLIP_OUTPUT = "Loaded text encoder for downstream CLIP inputs."
|
||||
VAE_OUTPUT = "Loaded VAE used to encode images to latents and decode latents to images."
|
||||
|
||||
VAE_OPTIONS_USE_TILING = (
|
||||
"Use ComfyUI's tiled VAE node. Disabled uses normal ComfyUI VAE behavior, "
|
||||
"including its automatic tiled retry after out-of-memory."
|
||||
)
|
||||
VAE_OPTIONS_ENCODE_PIXELS = "Image to encode into latent space."
|
||||
VAE_OPTIONS_DECODE_SAMPLES = "Latent samples to decode into an image."
|
||||
VAE_OPTIONS_VAE = "VAE used for the selected encode or decode operation."
|
||||
VAE_OPTIONS_TILE_SIZE = (
|
||||
"Tile size in pixels. Larger tiles are faster but use more memory."
|
||||
)
|
||||
VAE_OPTIONS_OVERLAP = (
|
||||
"Overlap between tiles in pixels. Larger overlaps reduce seams but do more work."
|
||||
)
|
||||
VAE_OPTIONS_ENCODE_TEMPORAL_SIZE = "For video VAEs, number of frames to encode at once."
|
||||
VAE_OPTIONS_DECODE_TEMPORAL_SIZE = "For video VAEs, number of frames to decode at once."
|
||||
VAE_OPTIONS_TEMPORAL_OVERLAP = (
|
||||
"For video VAEs, number of overlapping frames between temporal tiles."
|
||||
)
|
||||
VAE_OPTIONS_LATENT_OUTPUT = "Latent produced by ComfyUI's selected VAE encode node."
|
||||
VAE_OPTIONS_IMAGE_OUTPUT = "Image produced by ComfyUI's selected VAE decode node."
|
||||
|
||||
SAM_MODEL_OUTPUT = "Loaded SAM model for prompt-based mask and SEGS creation."
|
||||
GROUNDING_DINO_MODEL_OUTPUT = (
|
||||
"Loaded GroundingDINO model for finding prompt-matched boxes in images."
|
||||
@@ -40,6 +61,55 @@ GROUNDING_DINO_TEXT_ENCODER_INPUT = (
|
||||
VITMATTE_MODEL_INPUT = "ViTMatte model choice used for mask edge refinement."
|
||||
WD14_MODEL_INPUT = "WD14 tagger model choice used to generate tags from image crops."
|
||||
|
||||
EXTERNAL_LLM_MODEL_INPUT = "External model used for the prompt request."
|
||||
EXTERNAL_LLM_SYSTEM_PROMPT_INPUT = "Instruction text sent as the system message."
|
||||
EXTERNAL_LLM_USER_PROMPT_INPUT = "Prompt text sent as the user message."
|
||||
EXTERNAL_LLM_MAX_TOKENS_INPUT = (
|
||||
"Maximum number of response tokens the external model may generate. Higher "
|
||||
"values allow longer replies but can take longer and cost more."
|
||||
)
|
||||
EXTERNAL_LLM_REASONING_EFFORT_INPUT = (
|
||||
"Reasoning behavior for compatible providers. Default omits provider-specific "
|
||||
"controls; off sends thinking disabled through chat template options."
|
||||
)
|
||||
EXTERNAL_LLM_IMAGE_INPUT = (
|
||||
"Optional image sent with the user message for vision-capable external models. "
|
||||
"When a batch is connected, the first image is used."
|
||||
)
|
||||
EXTERNAL_LLM_RESPONSE_OUTPUT = "Assistant response returned by the external model."
|
||||
|
||||
EXTERNAL_LLM_TAG_SEGS_IMAGE = "Source image that the incoming SEGS were detected from."
|
||||
EXTERNAL_LLM_TAG_SEGS_SEGS = (
|
||||
"Existing SEGS to describe and keep aligned with the conditioning batch."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_CLIP = "CLIP model used to encode each generated regional prompt."
|
||||
EXTERNAL_LLM_TAG_SEGS_MODEL = "External vision model used to describe each SEG crop."
|
||||
EXTERNAL_LLM_TAG_SEGS_SYSTEM_PROMPT = (
|
||||
"Instruction text that controls how the model writes regional tags."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_USER_PROMPT = "Per-region request sent with each SEG crop image."
|
||||
EXTERNAL_LLM_TAG_SEGS_UNIVERSAL_POSITIVE = (
|
||||
"Positive prompt text added before every generated regional prompt."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_IMAGE_MODE = (
|
||||
"How pixels outside each SEG mask are shown to the vision model."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_REPLACE_UNDERSCORE = (
|
||||
"Replace underscores in generated tags before CLIP encoding."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_TRAILING_COMMA = (
|
||||
"Add a final comma to each generated regional prompt."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_EXCLUDE_TAGS = (
|
||||
"Comma-separated exact tags removed from generated regional prompts."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_SEGS_OUTPUT = (
|
||||
"Original SEGS returned in the same order as the conditioning batch."
|
||||
)
|
||||
EXTERNAL_LLM_TAG_SEGS_POSITIVE_OUTPUT = (
|
||||
"Positive conditioning from external LLM tags, matched to SEGS order."
|
||||
)
|
||||
|
||||
SAMPLING_MODEL = "Diffusion model used to denoise the input latent."
|
||||
SAMPLING_SEED = (
|
||||
"Seed used to create sampling noise. Reusing it with matching settings makes "
|
||||
@@ -188,3 +258,31 @@ TILE_SEGS_OUTPUT = "Generated tile SEGS in the same order as the conditioning ba
|
||||
TILE_POSITIVE_OUTPUT = (
|
||||
"Positive conditioning from WD14 tile tags, matched to SEGS order."
|
||||
)
|
||||
|
||||
TAG_SEGS_IMAGE = "Image that the incoming SEGS were detected from."
|
||||
TAG_SEGS_SEGS = "Existing SEGS to crop, tag, and keep in their current order."
|
||||
TAG_SEGS_CLIP = "CLIP model used to encode each generated SEGS prompt."
|
||||
TAG_SEGS_WD14_TAGGER = "WD14 tagger that reads each SEG crop and suggests prompt tags."
|
||||
TAG_SEGS_UNIVERSAL_POSITIVE = (
|
||||
"Positive prompt text added before every generated SEGS tag prompt."
|
||||
)
|
||||
TAG_SEGS_SEGS_OUTPUT = (
|
||||
"Original SEGS returned in the same order as the conditioning batch."
|
||||
)
|
||||
TAG_SEGS_POSITIVE_OUTPUT = (
|
||||
"Positive conditioning from WD14 SEGS tags, matched to SEGS order."
|
||||
)
|
||||
|
||||
BATCH_SEGS_FIRST = "First SEGS payload in the output order."
|
||||
BATCH_SEGS_SECOND = "Second SEGS payload appended after the first."
|
||||
BATCH_SEGS_OUTPUT = "Combined SEGS with all input segments in order."
|
||||
|
||||
BATCH_REGION_CONDITIONING_FIRST = (
|
||||
"First conditioning or conditioning batch in the output order."
|
||||
)
|
||||
BATCH_REGION_CONDITIONING_SECOND = (
|
||||
"Second conditioning or conditioning batch appended after the first."
|
||||
)
|
||||
BATCH_REGION_CONDITIONING_OUTPUT = (
|
||||
"Conditioning batch containing all input entries in order."
|
||||
)
|
||||
|
||||
@@ -0,0 +1,248 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""ComfyUI nodes that switch between native normal and tiled VAE execution."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..runtime.vae_options_graph import ExpansionResult, VAEOptionsGraphBuilder
|
||||
from . import tooltips
|
||||
|
||||
VAE_TILE_SIZE_DEFAULT = 512
|
||||
VAE_TILE_SIZE_MIN = 64
|
||||
VAE_TILE_SIZE_MAX = 4096
|
||||
VAE_ENCODE_TILE_SIZE_STEP = 64
|
||||
VAE_DECODE_TILE_SIZE_STEP = 32
|
||||
VAE_OVERLAP_DEFAULT = 64
|
||||
VAE_OVERLAP_MIN = 0
|
||||
VAE_OVERLAP_MAX = 4096
|
||||
VAE_OVERLAP_STEP = 32
|
||||
VAE_TEMPORAL_SIZE_DEFAULT = 64
|
||||
VAE_TEMPORAL_SIZE_MIN = 8
|
||||
VAE_TEMPORAL_SIZE_MAX = 4096
|
||||
VAE_TEMPORAL_SIZE_STEP = 4
|
||||
VAE_TEMPORAL_OVERLAP_DEFAULT = 8
|
||||
VAE_TEMPORAL_OVERLAP_MIN = 4
|
||||
VAE_TEMPORAL_OVERLAP_MAX = 4096
|
||||
VAE_TEMPORAL_OVERLAP_STEP = 4
|
||||
|
||||
|
||||
class VAEEncodeOptions:
|
||||
"""Encode images through ComfyUI's normal or tiled VAE encode nodes."""
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
RETURN_NAMES = ("latent",)
|
||||
OUTPUT_TOOLTIPS = (tooltips.VAE_OPTIONS_LATENT_OUTPUT,)
|
||||
FUNCTION = "encode"
|
||||
CATEGORY = "SimpleSyrup/Latent"
|
||||
DESCRIPTION = (
|
||||
"Encodes images to latent space with selectable normal or tiled VAE execution."
|
||||
)
|
||||
SEARCH_ALIASES = [
|
||||
"vae encode",
|
||||
"encode image",
|
||||
"tiled vae encode",
|
||||
"image to latent",
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare VAE encode inputs and tiled execution controls."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"use_tiling": _use_tiling_input(),
|
||||
"pixels": (
|
||||
"IMAGE",
|
||||
{
|
||||
"rawLink": True,
|
||||
"tooltip": tooltips.VAE_OPTIONS_ENCODE_PIXELS,
|
||||
},
|
||||
),
|
||||
"vae": _vae_input(),
|
||||
"tile_size": _tile_size_input(VAE_ENCODE_TILE_SIZE_STEP),
|
||||
"overlap": _overlap_input(),
|
||||
"temporal_size": _temporal_size_input(
|
||||
tooltips.VAE_OPTIONS_ENCODE_TEMPORAL_SIZE
|
||||
),
|
||||
"temporal_overlap": _temporal_overlap_input(),
|
||||
},
|
||||
}
|
||||
|
||||
def encode(
|
||||
self,
|
||||
use_tiling: bool,
|
||||
pixels: object,
|
||||
vae: object,
|
||||
tile_size: int,
|
||||
overlap: int,
|
||||
temporal_size: int = VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
temporal_overlap: int = VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
) -> ExpansionResult:
|
||||
"""Expand to ComfyUI's selected native VAE encode node."""
|
||||
|
||||
return VAEOptionsGraphBuilder().build_encode(
|
||||
pixels=pixels,
|
||||
vae=vae,
|
||||
use_tiling=use_tiling,
|
||||
tile_size=tile_size,
|
||||
overlap=overlap,
|
||||
temporal_size=temporal_size,
|
||||
temporal_overlap=temporal_overlap,
|
||||
)
|
||||
|
||||
|
||||
class VAEDecodeOptions:
|
||||
"""Decode latents through ComfyUI's normal or tiled VAE decode nodes."""
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
OUTPUT_TOOLTIPS = (tooltips.VAE_OPTIONS_IMAGE_OUTPUT,)
|
||||
FUNCTION = "decode"
|
||||
CATEGORY = "SimpleSyrup/Latent"
|
||||
DESCRIPTION = (
|
||||
"Decodes latents to images with selectable normal or tiled VAE execution."
|
||||
)
|
||||
SEARCH_ALIASES = [
|
||||
"vae decode",
|
||||
"decode latent",
|
||||
"tiled vae decode",
|
||||
"latent to image",
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Declare VAE decode inputs and tiled execution controls."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"use_tiling": _use_tiling_input(),
|
||||
"samples": (
|
||||
"LATENT",
|
||||
{
|
||||
"rawLink": True,
|
||||
"tooltip": tooltips.VAE_OPTIONS_DECODE_SAMPLES,
|
||||
},
|
||||
),
|
||||
"vae": _vae_input(),
|
||||
"tile_size": _tile_size_input(VAE_DECODE_TILE_SIZE_STEP),
|
||||
"overlap": _overlap_input(),
|
||||
"temporal_size": _temporal_size_input(
|
||||
tooltips.VAE_OPTIONS_DECODE_TEMPORAL_SIZE
|
||||
),
|
||||
"temporal_overlap": _temporal_overlap_input(),
|
||||
},
|
||||
}
|
||||
|
||||
def decode(
|
||||
self,
|
||||
use_tiling: bool,
|
||||
samples: object,
|
||||
vae: object,
|
||||
tile_size: int,
|
||||
overlap: int,
|
||||
temporal_size: int = VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
temporal_overlap: int = VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
) -> ExpansionResult:
|
||||
"""Expand to ComfyUI's selected native VAE decode node."""
|
||||
|
||||
return VAEOptionsGraphBuilder().build_decode(
|
||||
samples=samples,
|
||||
vae=vae,
|
||||
use_tiling=use_tiling,
|
||||
tile_size=tile_size,
|
||||
overlap=overlap,
|
||||
temporal_size=temporal_size,
|
||||
temporal_overlap=temporal_overlap,
|
||||
)
|
||||
|
||||
|
||||
def _use_tiling_input() -> tuple[str, dict[str, object]]:
|
||||
"""Return the shared tiling toggle declaration."""
|
||||
|
||||
return (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": tooltips.VAE_OPTIONS_USE_TILING,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _vae_input() -> tuple[str, dict[str, object]]:
|
||||
"""Return the shared raw-link VAE input declaration."""
|
||||
|
||||
return (
|
||||
"VAE",
|
||||
{
|
||||
"rawLink": True,
|
||||
"tooltip": tooltips.VAE_OPTIONS_VAE,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _tile_size_input(step: int) -> tuple[str, dict[str, object]]:
|
||||
"""Return the tile-size input declaration for encode or decode."""
|
||||
|
||||
return (
|
||||
"INT",
|
||||
{
|
||||
"default": VAE_TILE_SIZE_DEFAULT,
|
||||
"min": VAE_TILE_SIZE_MIN,
|
||||
"max": VAE_TILE_SIZE_MAX,
|
||||
"step": step,
|
||||
"advanced": True,
|
||||
"tooltip": tooltips.VAE_OPTIONS_TILE_SIZE,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _overlap_input() -> tuple[str, dict[str, object]]:
|
||||
"""Return the shared tile-overlap input declaration."""
|
||||
|
||||
return (
|
||||
"INT",
|
||||
{
|
||||
"default": VAE_OVERLAP_DEFAULT,
|
||||
"min": VAE_OVERLAP_MIN,
|
||||
"max": VAE_OVERLAP_MAX,
|
||||
"step": VAE_OVERLAP_STEP,
|
||||
"advanced": True,
|
||||
"tooltip": tooltips.VAE_OPTIONS_OVERLAP,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _temporal_size_input(tooltip: str) -> tuple[str, dict[str, object]]:
|
||||
"""Return the shared temporal tile-size input declaration."""
|
||||
|
||||
return (
|
||||
"INT",
|
||||
{
|
||||
"default": VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
"min": VAE_TEMPORAL_SIZE_MIN,
|
||||
"max": VAE_TEMPORAL_SIZE_MAX,
|
||||
"step": VAE_TEMPORAL_SIZE_STEP,
|
||||
"advanced": True,
|
||||
"tooltip": tooltip,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _temporal_overlap_input() -> tuple[str, dict[str, object]]:
|
||||
"""Return the shared temporal overlap input declaration."""
|
||||
|
||||
return (
|
||||
"INT",
|
||||
{
|
||||
"default": VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
"min": VAE_TEMPORAL_OVERLAP_MIN,
|
||||
"max": VAE_TEMPORAL_OVERLAP_MAX,
|
||||
"step": VAE_TEMPORAL_OVERLAP_STEP,
|
||||
"advanced": True,
|
||||
"tooltip": tooltips.VAE_OPTIONS_TEMPORAL_OVERLAP,
|
||||
},
|
||||
)
|
||||
@@ -12,29 +12,96 @@ from ..runtime.prompt_control_availability import prompt_control_is_available
|
||||
def get_nodes() -> list[type[object]]:
|
||||
"""Return v3 nodes that can be advertised in this environment."""
|
||||
|
||||
from .batch_region_conditioning import BatchRegionConditioningV3
|
||||
from .batch_segs import BatchSEGSV3
|
||||
from .external_llm_prompt import ExternalLLMPromptV3
|
||||
from .legacy_node_wrappers import (
|
||||
ConditioningBatchAppendV3,
|
||||
ConditioningBatchStartV3,
|
||||
DetailSEGSAsRegionsV3,
|
||||
DetailSEGSByScaleFactorTiledDiffusionV3,
|
||||
DetailSEGSByScaleFactorV3,
|
||||
DetectSEGSWithUltralyticsV3,
|
||||
EncodePromptBatchV3,
|
||||
GroundedSAMModelInfoV3,
|
||||
GroundingDINOModelLoaderV3,
|
||||
KSamplerExtrasV3,
|
||||
KSamplerTiledDiffusionV3,
|
||||
LatentDiagnosticsV3,
|
||||
LayerStyleSAMModelsAdapterV3,
|
||||
LoadUltralyticsModelV3,
|
||||
PromptEncodeStyleAndNormalizationV3,
|
||||
PromptEncodeStyleV3,
|
||||
PromptSEGSWithSAMV3,
|
||||
ResizeImageToTargetV3,
|
||||
SAMModelLoaderV3,
|
||||
SeedV3,
|
||||
SimpleLoadAnimaV3,
|
||||
SimpleVAEEncodeV3,
|
||||
UpscaleLatentFromImageV3,
|
||||
ViTMatteModelLoaderV3,
|
||||
)
|
||||
from .scale_factor import ScaleFactorV3
|
||||
from .simple_load_checkpoint import SimpleLoadCheckpointV3
|
||||
from .tag_segs_with_external_llm import TagSEGSWithExternalLLMV3
|
||||
from .tag_segs_with_wd14 import TagSEGSWithWD14V3
|
||||
from .tile_and_tag_segs import TileAndTagSEGSV3
|
||||
from .vae_decode_options import VAEDecodeOptionsV3
|
||||
from .vae_encode_options import VAEEncodeOptionsV3
|
||||
from .wd14_tagger_loader import WD14TaggerLoaderV3
|
||||
|
||||
nodes: list[type[object]] = [
|
||||
BatchRegionConditioningV3,
|
||||
BatchSEGSV3,
|
||||
ConditioningBatchAppendV3,
|
||||
ConditioningBatchStartV3,
|
||||
DetailSEGSAsRegionsV3,
|
||||
DetailSEGSByScaleFactorTiledDiffusionV3,
|
||||
DetailSEGSByScaleFactorV3,
|
||||
DetectSEGSWithUltralyticsV3,
|
||||
EncodePromptBatchV3,
|
||||
ExternalLLMPromptV3,
|
||||
GroundedSAMModelInfoV3,
|
||||
GroundingDINOModelLoaderV3,
|
||||
KSamplerExtrasV3,
|
||||
KSamplerTiledDiffusionV3,
|
||||
LatentDiagnosticsV3,
|
||||
LayerStyleSAMModelsAdapterV3,
|
||||
LoadUltralyticsModelV3,
|
||||
PromptEncodeStyleAndNormalizationV3,
|
||||
PromptEncodeStyleV3,
|
||||
PromptSEGSWithSAMV3,
|
||||
ResizeImageToTargetV3,
|
||||
SAMModelLoaderV3,
|
||||
ScaleFactorV3,
|
||||
SeedV3,
|
||||
SimpleLoadAnimaV3,
|
||||
SimpleLoadCheckpointV3,
|
||||
SimpleVAEEncodeV3,
|
||||
TagSEGSWithExternalLLMV3,
|
||||
TagSEGSWithWD14V3,
|
||||
TileAndTagSEGSV3,
|
||||
UpscaleLatentFromImageV3,
|
||||
VAEDecodeOptionsV3,
|
||||
VAEEncodeOptionsV3,
|
||||
ViTMatteModelLoaderV3,
|
||||
WD14TaggerLoaderV3,
|
||||
]
|
||||
|
||||
if not prompt_control_is_available():
|
||||
return [
|
||||
WD14TaggerLoaderV3,
|
||||
TileAndTagSEGSV3,
|
||||
SimpleLoadCheckpointV3,
|
||||
ScaleFactorV3,
|
||||
]
|
||||
return nodes
|
||||
|
||||
from .encode_prompt_batch_with_prompt_control import (
|
||||
EncodePromptBatchWithPromptControl,
|
||||
)
|
||||
from .schedule_and_encode_prompts_with_prompt_control import (
|
||||
ScheduleAndEncodePromptsWithPromptControl,
|
||||
)
|
||||
|
||||
return [
|
||||
WD14TaggerLoaderV3,
|
||||
TileAndTagSEGSV3,
|
||||
SimpleLoadCheckpointV3,
|
||||
ScaleFactorV3,
|
||||
*nodes,
|
||||
EncodePromptBatchWithPromptControl,
|
||||
ScheduleAndEncodePromptsWithPromptControl,
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node wrapper for batching regional conditioning."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.conditioning_batch import batch_conditioning
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
ConditioningBatchIO: Any = (
|
||||
None if TYPE_CHECKING else _comfy_io.Custom("CONDITIONING_BATCH")
|
||||
)
|
||||
|
||||
|
||||
class BatchRegionConditioningV3(_ComfyNodeBase):
|
||||
"""Expose expandable mixed conditioning batching through Comfy's v3 API."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the Batch Region Conditioning v3 schema."""
|
||||
|
||||
conditioning_input = _comfy_io.MultiType.Input(
|
||||
"conditioning",
|
||||
[_comfy_io.Conditioning, ConditioningBatchIO],
|
||||
tooltip="Conditioning or conditioning batch to append in socket order.",
|
||||
)
|
||||
autogrow_template = _comfy_io.Autogrow.TemplatePrefix(
|
||||
conditioning_input,
|
||||
prefix="conditioning",
|
||||
min=2,
|
||||
max=50,
|
||||
)
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.BatchRegionConditioning",
|
||||
display_name="Batch Region Conditioning",
|
||||
category="SimpleSyrup/Conditioning",
|
||||
description=(
|
||||
"Combines CONDITIONING and CONDITIONING_BATCH inputs into one "
|
||||
"ordered regional conditioning batch."
|
||||
),
|
||||
search_aliases=[
|
||||
"batch",
|
||||
"conditioning batch",
|
||||
"region conditioning",
|
||||
"segs prompts",
|
||||
],
|
||||
inputs=[
|
||||
_comfy_io.Autogrow.Input(
|
||||
"conditioning_inputs",
|
||||
template=autogrow_template,
|
||||
tooltip=(
|
||||
"Expandable conditioning inputs flattened in socket order."
|
||||
),
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
ConditioningBatchIO.Output(
|
||||
"batch",
|
||||
tooltip=(
|
||||
"Conditioning batch containing all input entries in order."
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, conditioning_inputs: Mapping[str, object]) -> tuple[object]:
|
||||
"""Batch conditioning inputs in Autogrow order."""
|
||||
|
||||
return (batch_conditioning(tuple(conditioning_inputs.values())),)
|
||||
@@ -0,0 +1,71 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node wrapper for Batch SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.segs import batch_segs, to_impact_compatible_segs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class BatchSEGSV3(_ComfyNodeBase):
|
||||
"""Expose expandable SEGS batching through Comfy's v3 API."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the Batch SEGS v3 schema."""
|
||||
|
||||
autogrow_template = _comfy_io.Autogrow.TemplatePrefix(
|
||||
_comfy_io.SEGS.Input(
|
||||
"segs",
|
||||
tooltip="SEGS payload to append to the output batch.",
|
||||
),
|
||||
prefix="segs",
|
||||
min=2,
|
||||
max=50,
|
||||
)
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.BatchSEGS",
|
||||
display_name="Batch SEGS",
|
||||
category="SimpleSyrup/Detection",
|
||||
description="Combines multiple SEGS inputs into one ordered SEGS payload.",
|
||||
search_aliases=["batch", "merge", "join", "combine", "segs"],
|
||||
inputs=[
|
||||
_comfy_io.Autogrow.Input(
|
||||
"segs_inputs",
|
||||
template=autogrow_template,
|
||||
tooltip="Expandable SEGS inputs joined in socket order.",
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.SEGS.Output(
|
||||
"segs",
|
||||
tooltip="Combined SEGS with all input segments in order.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, segs_inputs: Mapping[str, object]) -> tuple[object]:
|
||||
"""Batch provided SEGS inputs in Autogrow order."""
|
||||
|
||||
native = batch_segs(segs_inputs.values())
|
||||
return (to_impact_compatible_segs(native),)
|
||||
@@ -0,0 +1,121 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node wrapper for external LLM prompt generation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.external_llm import (
|
||||
DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
EXTERNAL_LLM_REASONING_EFFORTS,
|
||||
)
|
||||
from ..nodes import tooltips
|
||||
from ..nodes.external_llm_prompt import (
|
||||
MAX_EXTERNAL_LLM_MAX_TOKENS,
|
||||
ExternalLLMPrompt,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class ExternalLLMPromptV3(_ComfyNodeBase):
|
||||
"""Expose External LLM Prompt through Comfy's v3 extension API."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the External LLM Prompt v3 schema."""
|
||||
|
||||
required = ExternalLLMPrompt.INPUT_TYPES()["required"]
|
||||
model_choices = list(required["model"][0])
|
||||
model_default = str(required["model"][1]["default"])
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.ExternalLLMPrompt",
|
||||
display_name="External LLM Prompt",
|
||||
category="SimpleSyrup/Prompting",
|
||||
description=(
|
||||
"Sends system and user prompts to a configured external LLM provider."
|
||||
),
|
||||
search_aliases=["llm", "openai", "prompt", "preprocess", "text"],
|
||||
inputs=[
|
||||
_comfy_io.Combo.Input(
|
||||
"model",
|
||||
options=model_choices,
|
||||
default=model_default,
|
||||
tooltip=tooltips.EXTERNAL_LLM_MODEL_INPUT,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"system_prompt",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=tooltips.EXTERNAL_LLM_SYSTEM_PROMPT_INPUT,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"user_prompt",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=tooltips.EXTERNAL_LLM_USER_PROMPT_INPUT,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"max_tokens",
|
||||
default=DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
min=1,
|
||||
max=MAX_EXTERNAL_LLM_MAX_TOKENS,
|
||||
step=1,
|
||||
tooltip=tooltips.EXTERNAL_LLM_MAX_TOKENS_INPUT,
|
||||
),
|
||||
_comfy_io.Combo.Input(
|
||||
"reasoning_effort",
|
||||
options=list(EXTERNAL_LLM_REASONING_EFFORTS),
|
||||
default=DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
tooltip=tooltips.EXTERNAL_LLM_REASONING_EFFORT_INPUT,
|
||||
),
|
||||
_comfy_io.Image.Input(
|
||||
"image",
|
||||
optional=True,
|
||||
tooltip=tooltips.EXTERNAL_LLM_IMAGE_INPUT,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.String.Output(
|
||||
"response",
|
||||
tooltip=tooltips.EXTERNAL_LLM_RESPONSE_OUTPUT,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
image: object | None = None,
|
||||
) -> tuple[str]:
|
||||
"""Run the legacy node implementation behind the v3 schema."""
|
||||
|
||||
return ExternalLLMPrompt().generate(
|
||||
model=model,
|
||||
system_prompt=system_prompt,
|
||||
user_prompt=user_prompt,
|
||||
max_tokens=max_tokens,
|
||||
reasoning_effort=reasoning_effort,
|
||||
image=image,
|
||||
)
|
||||
@@ -0,0 +1,513 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 wrappers for nodes whose behavior still lives in legacy modules."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes.conditioning_batch_pack import (
|
||||
ConditioningBatchAppend,
|
||||
ConditioningBatchStart,
|
||||
)
|
||||
from ..nodes.detail_segs_as_regions import DetailSEGSAsRegions
|
||||
from ..nodes.detail_segs_by_scale_factor import DetailSEGSByScaleFactor
|
||||
from ..nodes.detail_segs_by_scale_factor_tiled_diffusion import (
|
||||
DetailSEGSByScaleFactorTiledDiffusion,
|
||||
)
|
||||
from ..nodes.detect_segs_with_ultralytics import DetectSEGSWithUltralytics
|
||||
from ..nodes.encode_prompt_batch import EncodePromptBatch
|
||||
from ..nodes.grounded_sam_model_info import GroundedSAMModelInfo
|
||||
from ..nodes.grounding_dino_model_loader import GroundingDINOModelLoader
|
||||
from ..nodes.image_resize_to_target import ResizeImageToTarget
|
||||
from ..nodes.ksampler_extras import KSamplerExtras
|
||||
from ..nodes.ksampler_tiled_diffusion import KSamplerTiledDiffusion
|
||||
from ..nodes.latent_diagnostics import LatentDiagnostics
|
||||
from ..nodes.layerstyle_sam_models_adapter import LayerStyleSAMModelsAdapter
|
||||
from ..nodes.load_ultralytics_model import LoadUltralyticsModel
|
||||
from ..nodes.prompt_encode_style import PromptEncodeStyle
|
||||
from ..nodes.prompt_encode_style_and_normalization import (
|
||||
PromptEncodeStyleAndNormalization,
|
||||
)
|
||||
from ..nodes.prompt_segs_with_sam import PromptSEGSWithSAM
|
||||
from ..nodes.provenance_latent import SimpleVAEEncode, UpscaleLatentFromImage
|
||||
from ..nodes.sam_model_loader import SAMModelLoader
|
||||
from ..nodes.seed import Seed
|
||||
from ..nodes.simple_load_anima import SimpleLoadAnima
|
||||
from ..nodes.vitmatte_model_loader import ViTMatteModelLoader
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
hidden: ClassVar[Any]
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
_HIDDEN_INPUTS = {
|
||||
"PROMPT": "prompt",
|
||||
"DYNPROMPT": "dynprompt",
|
||||
"EXTRA_PNGINFO": "extra_pnginfo",
|
||||
"UNIQUE_ID": "unique_id",
|
||||
"AUTH_TOKEN_COMFY_ORG": "auth_token_comfy_org",
|
||||
"API_KEY_COMFY_ORG": "api_key_comfy_org",
|
||||
}
|
||||
|
||||
|
||||
class LegacyNodeV3Adapter(_ComfyNodeBase):
|
||||
"""Build a v3 schema and execution bridge for a legacy implementation class."""
|
||||
|
||||
LEGACY_NODE_CLASS: ClassVar[type[Any]]
|
||||
NODE_ID: ClassVar[str]
|
||||
DISPLAY_NAME: ClassVar[str]
|
||||
ENABLE_EXPAND: ClassVar[bool] = False
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare a v3 schema from the implementation class contract."""
|
||||
|
||||
legacy = cls.LEGACY_NODE_CLASS
|
||||
return _comfy_io.Schema(
|
||||
node_id=cls.NODE_ID,
|
||||
display_name=cls.DISPLAY_NAME,
|
||||
category=str(getattr(legacy, "CATEGORY", "SimpleSyrup")),
|
||||
description=str(getattr(legacy, "DESCRIPTION", "")),
|
||||
search_aliases=list(getattr(legacy, "SEARCH_ALIASES", [])),
|
||||
inputs=_v3_inputs(legacy.INPUT_TYPES()),
|
||||
outputs=_v3_outputs(legacy),
|
||||
hidden=_v3_hidden_inputs(legacy.INPUT_TYPES()),
|
||||
is_input_list=bool(getattr(legacy, "INPUT_IS_LIST", False)),
|
||||
is_output_node=bool(getattr(legacy, "OUTPUT_NODE", False)),
|
||||
enable_expand=cls.ENABLE_EXPAND,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, **kwargs: object) -> Any:
|
||||
"""Run the wrapped implementation with v3-provided inputs."""
|
||||
|
||||
values = dict(kwargs)
|
||||
for name, hidden_attr in _legacy_hidden_inputs(
|
||||
cls.LEGACY_NODE_CLASS.INPUT_TYPES()
|
||||
).items():
|
||||
if name not in values:
|
||||
values[name] = getattr(cls.hidden, hidden_attr)
|
||||
|
||||
function_name = str(cls.LEGACY_NODE_CLASS.FUNCTION)
|
||||
implementation = cls.LEGACY_NODE_CLASS()
|
||||
function = getattr(implementation, function_name)
|
||||
return function(**values)
|
||||
|
||||
|
||||
class ConditioningBatchStartV3(LegacyNodeV3Adapter):
|
||||
"""Expose Conditioning Batch Start through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = ConditioningBatchStart
|
||||
NODE_ID = "SimpleSyrup.ConditioningBatchStart"
|
||||
DISPLAY_NAME = "Conditioning Batch Start"
|
||||
|
||||
|
||||
class ConditioningBatchAppendV3(LegacyNodeV3Adapter):
|
||||
"""Expose Conditioning Batch Append through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = ConditioningBatchAppend
|
||||
NODE_ID = "SimpleSyrup.ConditioningBatchAppend"
|
||||
DISPLAY_NAME = "Conditioning Batch Append"
|
||||
|
||||
|
||||
class GroundedSAMModelInfoV3(LegacyNodeV3Adapter):
|
||||
"""Expose Grounded SAM Model Info through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = GroundedSAMModelInfo
|
||||
NODE_ID = "SimpleSyrup.GroundedSAMModelInfo"
|
||||
DISPLAY_NAME = "Grounded SAM Model Info"
|
||||
|
||||
|
||||
class GroundingDINOModelLoaderV3(LegacyNodeV3Adapter):
|
||||
"""Expose GroundingDINO Model Loader through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = GroundingDINOModelLoader
|
||||
NODE_ID = "SimpleSyrup.GroundingDINOModelLoader"
|
||||
DISPLAY_NAME = "GroundingDINO Model Loader"
|
||||
|
||||
|
||||
class KSamplerExtrasV3(LegacyNodeV3Adapter):
|
||||
"""Expose KSampler Extras through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = KSamplerExtras
|
||||
NODE_ID = "SimpleSyrup.KSamplerExtras"
|
||||
DISPLAY_NAME = "KSampler (Extras)"
|
||||
|
||||
|
||||
class KSamplerTiledDiffusionV3(LegacyNodeV3Adapter):
|
||||
"""Expose KSampler Tiled Diffusion through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = KSamplerTiledDiffusion
|
||||
NODE_ID = "SimpleSyrup.KSamplerTiledDiffusion"
|
||||
DISPLAY_NAME = "KSampler (Tiled Diffusion)"
|
||||
|
||||
|
||||
class LayerStyleSAMModelsAdapterV3(LegacyNodeV3Adapter):
|
||||
"""Expose LayerStyle SAM Models Adapter through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = LayerStyleSAMModelsAdapter
|
||||
NODE_ID = "SimpleSyrup.LayerStyleSAMModelsAdapter"
|
||||
DISPLAY_NAME = "LayerStyle SAM Models Adapter"
|
||||
|
||||
|
||||
class LatentDiagnosticsV3(LegacyNodeV3Adapter):
|
||||
"""Expose Latent Diagnostics through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = LatentDiagnostics
|
||||
NODE_ID = "SimpleSyrup.LatentDiagnostics"
|
||||
DISPLAY_NAME = "Latent Diagnostics"
|
||||
|
||||
|
||||
class PromptEncodeStyleV3(LegacyNodeV3Adapter):
|
||||
"""Expose Prompt Encode Style through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = PromptEncodeStyle
|
||||
NODE_ID = "SimpleSyrup.PromptEncodeStyle"
|
||||
DISPLAY_NAME = "Prompt Encode Style"
|
||||
|
||||
|
||||
class PromptEncodeStyleAndNormalizationV3(LegacyNodeV3Adapter):
|
||||
"""Expose Prompt Encode Style and Normalization through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = PromptEncodeStyleAndNormalization
|
||||
NODE_ID = "SimpleSyrup.PromptEncodeStyleAndNormalization"
|
||||
DISPLAY_NAME = "Prompt Encode Style & Normalization"
|
||||
|
||||
|
||||
class PromptSEGSWithSAMV3(LegacyNodeV3Adapter):
|
||||
"""Expose Prompt SEGS w/ SAM through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = PromptSEGSWithSAM
|
||||
NODE_ID = "SimpleSyrup.PromptSEGSWithSAM"
|
||||
DISPLAY_NAME = "Prompt SEGS w/ SAM"
|
||||
|
||||
|
||||
class SimpleVAEEncodeV3(LegacyNodeV3Adapter):
|
||||
"""Expose Simple VAE Encode through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = SimpleVAEEncode
|
||||
NODE_ID = "SimpleSyrup.SimpleVAEEncode"
|
||||
DISPLAY_NAME = "Simple VAE Encode"
|
||||
ENABLE_EXPAND = True
|
||||
|
||||
|
||||
class UpscaleLatentFromImageV3(LegacyNodeV3Adapter):
|
||||
"""Expose Upscale Latent From Image through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = UpscaleLatentFromImage
|
||||
NODE_ID = "SimpleSyrup.UpscaleLatentFromImage"
|
||||
DISPLAY_NAME = "Upscale Latent From Image"
|
||||
ENABLE_EXPAND = True
|
||||
|
||||
|
||||
class ResizeImageToTargetV3(LegacyNodeV3Adapter):
|
||||
"""Expose Resize Image to Target through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = ResizeImageToTarget
|
||||
NODE_ID = "SimpleSyrup.ResizeImageToTarget"
|
||||
DISPLAY_NAME = "Resize Image to Target"
|
||||
|
||||
|
||||
class DetailSEGSAsRegionsV3(LegacyNodeV3Adapter):
|
||||
"""Expose Detail SEGS as Regions through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = DetailSEGSAsRegions
|
||||
NODE_ID = "SimpleSyrup.DetailSEGSAsRegions"
|
||||
DISPLAY_NAME = "Detail SEGS as Regions"
|
||||
|
||||
|
||||
class DetailSEGSByScaleFactorV3(LegacyNodeV3Adapter):
|
||||
"""Expose Detail SEGS by Scale Factor through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = DetailSEGSByScaleFactor
|
||||
NODE_ID = "SimpleSyrup.DetailSEGSByScaleFactor"
|
||||
DISPLAY_NAME = "Detail SEGS by Scale Factor"
|
||||
|
||||
|
||||
class DetailSEGSByScaleFactorTiledDiffusionV3(LegacyNodeV3Adapter):
|
||||
"""Expose Detail SEGS by Scale Factor with Tiled Diffusion through Comfy v3."""
|
||||
|
||||
LEGACY_NODE_CLASS = DetailSEGSByScaleFactorTiledDiffusion
|
||||
NODE_ID = "SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion"
|
||||
DISPLAY_NAME = "Detail SEGS by Scale Factor w/ Tiled Diffusion"
|
||||
|
||||
|
||||
class SAMModelLoaderV3(LegacyNodeV3Adapter):
|
||||
"""Expose SAM Model Loader through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = SAMModelLoader
|
||||
NODE_ID = "SimpleSyrup.SAMModelLoader"
|
||||
DISPLAY_NAME = "SAM Model Loader"
|
||||
|
||||
|
||||
class SeedV3(LegacyNodeV3Adapter):
|
||||
"""Expose Seed through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = Seed
|
||||
NODE_ID = "SimpleSyrup.Seed"
|
||||
DISPLAY_NAME = "Seed"
|
||||
|
||||
|
||||
class SimpleLoadAnimaV3(LegacyNodeV3Adapter):
|
||||
"""Expose Simple Load Anima through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = SimpleLoadAnima
|
||||
NODE_ID = "SimpleSyrup.SimpleLoadAnima"
|
||||
DISPLAY_NAME = "Simple Load Anima"
|
||||
|
||||
|
||||
class LoadUltralyticsModelV3(LegacyNodeV3Adapter):
|
||||
"""Expose Load Ultralytics Model through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = LoadUltralyticsModel
|
||||
NODE_ID = "SimpleSyrup.LoadUltralyticsModel"
|
||||
DISPLAY_NAME = "Load Ultralytics Model"
|
||||
|
||||
|
||||
class DetectSEGSWithUltralyticsV3(LegacyNodeV3Adapter):
|
||||
"""Expose Detect SEGS w/ Ultralytics through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = DetectSEGSWithUltralytics
|
||||
NODE_ID = "SimpleSyrup.DetectSEGSWithUltralytics"
|
||||
DISPLAY_NAME = "Detect SEGS w/ Ultralytics"
|
||||
|
||||
|
||||
class EncodePromptBatchV3(LegacyNodeV3Adapter):
|
||||
"""Expose Encode Prompt Batch through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = EncodePromptBatch
|
||||
NODE_ID = "SimpleSyrup.EncodePromptBatch"
|
||||
DISPLAY_NAME = "Encode Prompt Batch"
|
||||
|
||||
|
||||
class ViTMatteModelLoaderV3(LegacyNodeV3Adapter):
|
||||
"""Expose ViTMatte Model Loader through Comfy v3 only."""
|
||||
|
||||
LEGACY_NODE_CLASS = ViTMatteModelLoader
|
||||
NODE_ID = "SimpleSyrup.ViTMatteModelLoader"
|
||||
DISPLAY_NAME = "ViTMatte Model Loader"
|
||||
|
||||
|
||||
def _v3_inputs(input_types: Mapping[str, Mapping[str, object]]) -> list[Any]:
|
||||
"""Return v3 input declarations from legacy required and optional inputs."""
|
||||
|
||||
inputs: list[Any] = []
|
||||
for section_name, optional in (("required", False), ("optional", True)):
|
||||
section = input_types.get(section_name, {})
|
||||
for name, declaration in section.items():
|
||||
inputs.append(_v3_input(name, declaration, optional=optional))
|
||||
return inputs
|
||||
|
||||
|
||||
def _v3_input(name: str, declaration: object, *, optional: bool) -> Any:
|
||||
"""Return one v3 input declaration from a legacy field declaration."""
|
||||
|
||||
if not isinstance(declaration, tuple) or not declaration:
|
||||
raise TypeError(f"legacy input {name} declaration must be a tuple.")
|
||||
|
||||
io_declaration = declaration[0]
|
||||
options = _input_options(declaration)
|
||||
tooltip = _string_option(options, "tooltip")
|
||||
advanced = _bool_option(options, "advanced")
|
||||
raw_link = _bool_option(options, "rawLink") or _bool_option(options, "raw_link")
|
||||
force_input = _bool_option(options, "forceInput") or _bool_option(
|
||||
options, "force_input"
|
||||
)
|
||||
|
||||
if isinstance(io_declaration, (list, tuple)):
|
||||
return _comfy_io.Combo.Input(
|
||||
name,
|
||||
options=list(io_declaration),
|
||||
optional=optional,
|
||||
default=options.get("default"),
|
||||
control_after_generate=options.get("control_after_generate"),
|
||||
tooltip=tooltip,
|
||||
raw_link=raw_link,
|
||||
advanced=advanced,
|
||||
)
|
||||
|
||||
if not isinstance(io_declaration, str):
|
||||
raise TypeError(f"legacy input {name} type must be a string or options list.")
|
||||
|
||||
input_type = io_declaration
|
||||
input_class = _io_class(input_type)
|
||||
common_options = {
|
||||
"optional": optional,
|
||||
"tooltip": tooltip,
|
||||
"raw_link": raw_link,
|
||||
"advanced": advanced,
|
||||
}
|
||||
|
||||
if input_type == "INT":
|
||||
return input_class.Input(
|
||||
name,
|
||||
default=options.get("default"),
|
||||
min=options.get("min"),
|
||||
max=options.get("max"),
|
||||
step=options.get("step"),
|
||||
control_after_generate=options.get("control_after_generate"),
|
||||
**common_options,
|
||||
)
|
||||
if input_type == "FLOAT":
|
||||
return input_class.Input(
|
||||
name,
|
||||
default=options.get("default"),
|
||||
min=options.get("min"),
|
||||
max=options.get("max"),
|
||||
step=options.get("step"),
|
||||
round=options.get("round"),
|
||||
**common_options,
|
||||
)
|
||||
if input_type == "STRING":
|
||||
return input_class.Input(
|
||||
name,
|
||||
default=options.get("default"),
|
||||
multiline=bool(options.get("multiline", False)),
|
||||
force_input=force_input,
|
||||
**common_options,
|
||||
)
|
||||
if input_type == "BOOLEAN":
|
||||
return input_class.Input(
|
||||
name,
|
||||
default=options.get("default"),
|
||||
label_on=options.get("label_on"),
|
||||
label_off=options.get("label_off"),
|
||||
**common_options,
|
||||
)
|
||||
|
||||
return input_class.Input(name, **common_options)
|
||||
|
||||
|
||||
def _v3_outputs(legacy: type[Any]) -> list[Any]:
|
||||
"""Return v3 output declarations from legacy return metadata."""
|
||||
|
||||
return_types = tuple(getattr(legacy, "RETURN_TYPES", ()))
|
||||
return_names = getattr(legacy, "RETURN_NAMES", None)
|
||||
output_tooltips = tuple(getattr(legacy, "OUTPUT_TOOLTIPS", ()))
|
||||
output_is_list = tuple(
|
||||
getattr(legacy, "OUTPUT_IS_LIST", (False,) * len(return_types))
|
||||
)
|
||||
outputs: list[Any] = []
|
||||
for index, io_type in enumerate(return_types):
|
||||
output_name = None
|
||||
if isinstance(return_names, tuple) and index < len(return_names):
|
||||
output_name = str(return_names[index])
|
||||
tooltip = None
|
||||
if index < len(output_tooltips):
|
||||
tooltip = str(output_tooltips[index])
|
||||
is_output_list = index < len(output_is_list) and bool(output_is_list[index])
|
||||
outputs.append(
|
||||
_io_class(str(io_type)).Output(
|
||||
output_name,
|
||||
tooltip=tooltip,
|
||||
is_output_list=is_output_list,
|
||||
)
|
||||
)
|
||||
return outputs
|
||||
|
||||
|
||||
def _v3_hidden_inputs(input_types: Mapping[str, Mapping[str, object]]) -> list[Any]:
|
||||
"""Return v3 hidden declarations requested by legacy hidden inputs."""
|
||||
|
||||
hidden_values = set(_legacy_hidden_inputs(input_types).values())
|
||||
return [getattr(_comfy_io.Hidden, value) for value in sorted(hidden_values)]
|
||||
|
||||
|
||||
def _legacy_hidden_inputs(
|
||||
input_types: Mapping[str, Mapping[str, object]],
|
||||
) -> dict[str, str]:
|
||||
"""Return legacy hidden input names mapped to v3 hidden holder attributes."""
|
||||
|
||||
hidden_inputs: dict[str, str] = {}
|
||||
for name, sentinel in input_types.get("hidden", {}).items():
|
||||
if isinstance(sentinel, str) and sentinel in _HIDDEN_INPUTS:
|
||||
hidden_inputs[name] = _HIDDEN_INPUTS[sentinel]
|
||||
return hidden_inputs
|
||||
|
||||
|
||||
def _io_class(io_type: str) -> Any:
|
||||
"""Return the v3 IO class for a legacy Comfy type string."""
|
||||
|
||||
known_types = {
|
||||
"BOOLEAN": _comfy_io.Boolean,
|
||||
"INT": _comfy_io.Int,
|
||||
"FLOAT": _comfy_io.Float,
|
||||
"STRING": _comfy_io.String,
|
||||
"IMAGE": _comfy_io.Image,
|
||||
"MASK": _comfy_io.Mask,
|
||||
"LATENT": _comfy_io.Latent,
|
||||
"MODEL": _comfy_io.Model,
|
||||
"CLIP": _comfy_io.Clip,
|
||||
"VAE": _comfy_io.Vae,
|
||||
"CONDITIONING": _comfy_io.Conditioning,
|
||||
"SEGS": _comfy_io.SEGS,
|
||||
}
|
||||
return known_types.get(io_type, _comfy_io.Custom(io_type))
|
||||
|
||||
|
||||
def _input_options(declaration: tuple[object, ...]) -> dict[str, object]:
|
||||
"""Return an input options dictionary from a legacy declaration."""
|
||||
|
||||
if len(declaration) < 2 or not isinstance(declaration[1], dict):
|
||||
return {}
|
||||
return dict(declaration[1])
|
||||
|
||||
|
||||
def _string_option(options: Mapping[str, object], name: str) -> str | None:
|
||||
"""Return a string option when present."""
|
||||
|
||||
value = options.get(name)
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
def _bool_option(options: Mapping[str, object], name: str) -> bool | None:
|
||||
"""Return a boolean option when present."""
|
||||
|
||||
value = options.get(name)
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ConditioningBatchAppendV3",
|
||||
"ConditioningBatchStartV3",
|
||||
"DetailSEGSAsRegionsV3",
|
||||
"DetailSEGSByScaleFactorTiledDiffusionV3",
|
||||
"DetailSEGSByScaleFactorV3",
|
||||
"DetectSEGSWithUltralyticsV3",
|
||||
"EncodePromptBatchV3",
|
||||
"GroundedSAMModelInfoV3",
|
||||
"GroundingDINOModelLoaderV3",
|
||||
"KSamplerExtrasV3",
|
||||
"KSamplerTiledDiffusionV3",
|
||||
"LatentDiagnosticsV3",
|
||||
"LayerStyleSAMModelsAdapterV3",
|
||||
"LoadUltralyticsModelV3",
|
||||
"PromptEncodeStyleAndNormalizationV3",
|
||||
"PromptEncodeStyleV3",
|
||||
"PromptSEGSWithSAMV3",
|
||||
"ResizeImageToTargetV3",
|
||||
"SAMModelLoaderV3",
|
||||
"SeedV3",
|
||||
"SimpleLoadAnimaV3",
|
||||
"SimpleVAEEncodeV3",
|
||||
"UpscaleLatentFromImageV3",
|
||||
"ViTMatteModelLoaderV3",
|
||||
]
|
||||
@@ -0,0 +1,137 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for Prompt-Control prompt scheduling and encoding."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..runtime.prompt_control_schedule_encode_graph import (
|
||||
PromptControlScheduleEncodeGraphBuilder,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Return v1-compatible input metadata."""
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
MixedConditioningIO: Any = (
|
||||
None if TYPE_CHECKING else _comfy_io.Custom("CONDITIONING,CONDITIONING_BATCH")
|
||||
)
|
||||
|
||||
|
||||
class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase):
|
||||
"""Schedule Prompt-Control LoRAs and encode prompts with optional batches."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the Prompt-Control schedule and encode schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl",
|
||||
display_name="Schedule & Encode Prompts",
|
||||
enable_expand=True,
|
||||
category="SimpleSyrup/Conditioning",
|
||||
description=(
|
||||
"Schedules Prompt-Control LoRAs and encodes prompts, using [SEP] "
|
||||
"to create SimpleSyrup conditioning batches."
|
||||
),
|
||||
inputs=[
|
||||
_comfy_io.Model.Input(
|
||||
"model",
|
||||
raw_link=True,
|
||||
tooltip=(
|
||||
"Model that receives LoRA changes found in Prompt-Control "
|
||||
"prompt tags."
|
||||
),
|
||||
),
|
||||
_comfy_io.Clip.Input(
|
||||
"clip",
|
||||
raw_link=True,
|
||||
tooltip=(
|
||||
"CLIP connection used for scheduled hooks and cleaned "
|
||||
"prompt encoding."
|
||||
),
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"encode_style",
|
||||
default="",
|
||||
force_input=True,
|
||||
optional=True,
|
||||
tooltip=(
|
||||
"Encode style text from Prompt Encode Style or "
|
||||
"Prompt Encode Style & Normalization."
|
||||
),
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"positive_prompt",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=(
|
||||
"Positive Prompt-Control text. [SEP] creates a "
|
||||
"conditioning batch for SimpleSyrup batch-aware nodes."
|
||||
),
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"negative_prompt",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=(
|
||||
"Negative Prompt-Control text. [SEP] creates a "
|
||||
"conditioning batch for SimpleSyrup batch-aware nodes."
|
||||
),
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Model.Output(
|
||||
"model",
|
||||
tooltip=(
|
||||
"Model after LoRA tags from positive and negative prompts "
|
||||
"are scheduled."
|
||||
),
|
||||
),
|
||||
MixedConditioningIO.Output(
|
||||
"positive",
|
||||
tooltip="Positive conditioning or SimpleSyrup conditioning batch.",
|
||||
),
|
||||
MixedConditioningIO.Output(
|
||||
"negative",
|
||||
tooltip="Negative conditioning or SimpleSyrup conditioning batch.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
model: Any,
|
||||
clip: Any,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
encode_style: str = "",
|
||||
) -> Any:
|
||||
"""Build lazy Prompt-Control graph expansion for prompts."""
|
||||
|
||||
return PromptControlScheduleEncodeGraphBuilder().build(
|
||||
model=model,
|
||||
clip=clip,
|
||||
positive_prompt=positive_prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
encode_style=encode_style,
|
||||
)
|
||||
@@ -0,0 +1,193 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node for tagging existing SEGS with an external LLM."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..domain.external_llm import (
|
||||
DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
EXTERNAL_LLM_REASONING_EFFORTS,
|
||||
)
|
||||
from ..nodes import tooltips
|
||||
from ..runtime.external_llm_images import SEG_IMAGE_MODES
|
||||
from ..services.tag_segs_with_external_llm_service import (
|
||||
LLMTagFormattingControls,
|
||||
TagSEGSWithExternalLLMService,
|
||||
)
|
||||
|
||||
MAX_EXTERNAL_LLM_MAX_TOKENS = 32768
|
||||
DEFAULT_SEG_IMAGE_MODE = "transparent mask"
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
ConditioningBatchIO: Any = (
|
||||
None if TYPE_CHECKING else _comfy_io.Custom("CONDITIONING_BATCH")
|
||||
)
|
||||
|
||||
|
||||
class TagSEGSWithExternalLLMV3(_ComfyNodeBase):
|
||||
"""Expose external-LLM SEGS tagging through Comfy's v3 API."""
|
||||
|
||||
_service = TagSEGSWithExternalLLMService()
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the Tag SEGS w/ External LLM v3 schema."""
|
||||
|
||||
model_choices = cls._service.model_choices()
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.TagSEGSWithExternalLLM",
|
||||
display_name="Tag SEGS w/ External LLM",
|
||||
category="SimpleSyrup/Detailing",
|
||||
description=(
|
||||
"Tags existing SEGS crops with a configured external vision LLM "
|
||||
"and returns aligned conditioning for SEGS detailing."
|
||||
),
|
||||
search_aliases=[
|
||||
"llm",
|
||||
"vision",
|
||||
"tag",
|
||||
"segs",
|
||||
"detail",
|
||||
"regional",
|
||||
"prompt",
|
||||
],
|
||||
inputs=[
|
||||
_comfy_io.Image.Input(
|
||||
"image",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_IMAGE,
|
||||
),
|
||||
_comfy_io.SEGS.Input(
|
||||
"segs",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_SEGS,
|
||||
),
|
||||
_comfy_io.Clip.Input(
|
||||
"clip",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_CLIP,
|
||||
),
|
||||
_comfy_io.Combo.Input(
|
||||
"model",
|
||||
options=model_choices,
|
||||
default=model_choices[0],
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_MODEL,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"system_prompt",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_SYSTEM_PROMPT,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"user_prompt",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_USER_PROMPT,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"universal_positive",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_UNIVERSAL_POSITIVE,
|
||||
),
|
||||
_comfy_io.Combo.Input(
|
||||
"seg_image_mode",
|
||||
options=list(SEG_IMAGE_MODES),
|
||||
default=DEFAULT_SEG_IMAGE_MODE,
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_IMAGE_MODE,
|
||||
),
|
||||
_comfy_io.Boolean.Input(
|
||||
"replace_underscore",
|
||||
default=True,
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_REPLACE_UNDERSCORE,
|
||||
),
|
||||
_comfy_io.Boolean.Input(
|
||||
"trailing_comma",
|
||||
default=False,
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_TRAILING_COMMA,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"exclude_tags",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_EXCLUDE_TAGS,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"max_tokens",
|
||||
default=DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
min=1,
|
||||
max=MAX_EXTERNAL_LLM_MAX_TOKENS,
|
||||
step=1,
|
||||
tooltip=tooltips.EXTERNAL_LLM_MAX_TOKENS_INPUT,
|
||||
),
|
||||
_comfy_io.Combo.Input(
|
||||
"reasoning_effort",
|
||||
options=list(EXTERNAL_LLM_REASONING_EFFORTS),
|
||||
default=DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
tooltip=tooltips.EXTERNAL_LLM_REASONING_EFFORT_INPUT,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.SEGS.Output(
|
||||
"segs",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_SEGS_OUTPUT,
|
||||
),
|
||||
ConditioningBatchIO.Output(
|
||||
"positive",
|
||||
tooltip=tooltips.EXTERNAL_LLM_TAG_SEGS_POSITIVE_OUTPUT,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
image: object,
|
||||
segs: object,
|
||||
clip: Any,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
universal_positive: str,
|
||||
seg_image_mode: str,
|
||||
replace_underscore: bool,
|
||||
trailing_comma: bool,
|
||||
exclude_tags: str,
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
) -> tuple[object, object]:
|
||||
"""Tag existing SEGS and return aligned conditioning."""
|
||||
|
||||
result = cls._service.tag(
|
||||
image=image,
|
||||
segs=segs,
|
||||
clip=clip,
|
||||
model=model,
|
||||
system_prompt=system_prompt,
|
||||
user_prompt=user_prompt,
|
||||
universal_positive=universal_positive,
|
||||
seg_image_mode=seg_image_mode,
|
||||
formatting=LLMTagFormattingControls(
|
||||
replace_underscore=replace_underscore,
|
||||
trailing_comma=trailing_comma,
|
||||
exclude_tags=exclude_tags,
|
||||
),
|
||||
max_tokens=max_tokens,
|
||||
reasoning_effort=reasoning_effort,
|
||||
)
|
||||
return result.segs, result.positive
|
||||
@@ -0,0 +1,136 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node wrapper for Tag SEGS w/ WD14."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..nodes.tag_segs_with_wd14 import TagSEGSWithWD14
|
||||
from ..nodes.tile_and_tag_segs import DEFAULT_EXCLUDE_TAGS
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
ConditioningBatchIO: Any = (
|
||||
None if TYPE_CHECKING else _comfy_io.Custom("CONDITIONING_BATCH")
|
||||
)
|
||||
WD14TaggerIO: Any = None if TYPE_CHECKING else _comfy_io.Custom("WD14_TAGGER")
|
||||
|
||||
|
||||
class TagSEGSWithWD14V3(_ComfyNodeBase):
|
||||
"""Expose WD14 tagging for existing SEGS through Comfy's v3 API."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the Tag SEGS w/ WD14 v3 schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.TagSEGSWithWD14",
|
||||
display_name="Tag SEGS w/ WD14",
|
||||
category="SimpleSyrup/Detailing",
|
||||
description=(
|
||||
"Tags existing SEGS crops with a connected WD14 tagger and "
|
||||
"returns aligned conditioning for SEGS detailing."
|
||||
),
|
||||
search_aliases=["tag", "wd14", "segs", "detail", "regional"],
|
||||
inputs=[
|
||||
_comfy_io.Image.Input("image", tooltip=tooltips.TAG_SEGS_IMAGE),
|
||||
_comfy_io.SEGS.Input("segs", tooltip=tooltips.TAG_SEGS_SEGS),
|
||||
_comfy_io.Clip.Input("clip", tooltip=tooltips.TAG_SEGS_CLIP),
|
||||
WD14TaggerIO.Input(
|
||||
"wd14_tagger",
|
||||
tooltip=tooltips.TAG_SEGS_WD14_TAGGER,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"universal_positive",
|
||||
multiline=False,
|
||||
default="",
|
||||
tooltip=tooltips.TAG_SEGS_UNIVERSAL_POSITIVE,
|
||||
),
|
||||
_comfy_io.Float.Input(
|
||||
"threshold",
|
||||
default=0.35,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.05,
|
||||
tooltip=tooltips.TILE_THRESHOLD,
|
||||
),
|
||||
_comfy_io.Float.Input(
|
||||
"character_threshold",
|
||||
default=1.0,
|
||||
min=0.0,
|
||||
max=1.0,
|
||||
step=0.05,
|
||||
tooltip=tooltips.TILE_CHARACTER_THRESHOLD,
|
||||
),
|
||||
_comfy_io.Boolean.Input(
|
||||
"replace_underscore",
|
||||
default=True,
|
||||
tooltip=tooltips.TILE_REPLACE_UNDERSCORE,
|
||||
),
|
||||
_comfy_io.Boolean.Input(
|
||||
"trailing_comma",
|
||||
default=False,
|
||||
tooltip=tooltips.TILE_TRAILING_COMMA,
|
||||
),
|
||||
_comfy_io.String.Input(
|
||||
"exclude_tags",
|
||||
multiline=False,
|
||||
default=DEFAULT_EXCLUDE_TAGS,
|
||||
tooltip=tooltips.TILE_EXCLUDE_TAGS,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.SEGS.Output(
|
||||
"segs",
|
||||
tooltip=tooltips.TAG_SEGS_SEGS_OUTPUT,
|
||||
),
|
||||
ConditioningBatchIO.Output(
|
||||
"positive",
|
||||
tooltip=tooltips.TAG_SEGS_POSITIVE_OUTPUT,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
image: object,
|
||||
segs: object,
|
||||
clip: Any,
|
||||
wd14_tagger: object,
|
||||
universal_positive: str,
|
||||
threshold: float,
|
||||
character_threshold: float,
|
||||
replace_underscore: bool,
|
||||
trailing_comma: bool,
|
||||
exclude_tags: str,
|
||||
) -> tuple[object, object]:
|
||||
"""Run the legacy implementation behind the v3 schema."""
|
||||
|
||||
return TagSEGSWithWD14().tag(
|
||||
image=image,
|
||||
segs=segs,
|
||||
clip=clip,
|
||||
wd14_tagger=wd14_tagger,
|
||||
universal_positive=universal_positive,
|
||||
threshold=threshold,
|
||||
character_threshold=character_threshold,
|
||||
replace_underscore=replace_underscore,
|
||||
trailing_comma=trailing_comma,
|
||||
exclude_tags=exclude_tags,
|
||||
)
|
||||
@@ -0,0 +1,143 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node wrapper for VAE Decode (Options)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..nodes.vae_options import (
|
||||
VAE_DECODE_TILE_SIZE_STEP,
|
||||
VAE_OVERLAP_DEFAULT,
|
||||
VAE_OVERLAP_MAX,
|
||||
VAE_OVERLAP_MIN,
|
||||
VAE_OVERLAP_STEP,
|
||||
VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
VAE_TEMPORAL_OVERLAP_MAX,
|
||||
VAE_TEMPORAL_OVERLAP_MIN,
|
||||
VAE_TEMPORAL_OVERLAP_STEP,
|
||||
VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
VAE_TEMPORAL_SIZE_MAX,
|
||||
VAE_TEMPORAL_SIZE_MIN,
|
||||
VAE_TEMPORAL_SIZE_STEP,
|
||||
VAE_TILE_SIZE_DEFAULT,
|
||||
VAE_TILE_SIZE_MAX,
|
||||
VAE_TILE_SIZE_MIN,
|
||||
VAEDecodeOptions,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class VAEDecodeOptionsV3(_ComfyNodeBase):
|
||||
"""Expose VAE Decode (Options) through Comfy's v3 extension API."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the VAE Decode (Options) v3 schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.VAEDecodeOptions",
|
||||
display_name="VAE Decode (Options)",
|
||||
enable_expand=True,
|
||||
category="SimpleSyrup/Latent",
|
||||
description=VAEDecodeOptions.DESCRIPTION,
|
||||
search_aliases=VAEDecodeOptions.SEARCH_ALIASES,
|
||||
inputs=[
|
||||
_comfy_io.Boolean.Input(
|
||||
"use_tiling",
|
||||
default=False,
|
||||
tooltip=tooltips.VAE_OPTIONS_USE_TILING,
|
||||
),
|
||||
_comfy_io.Latent.Input(
|
||||
"samples",
|
||||
raw_link=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_DECODE_SAMPLES,
|
||||
),
|
||||
_comfy_io.Vae.Input(
|
||||
"vae",
|
||||
raw_link=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_VAE,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"tile_size",
|
||||
default=VAE_TILE_SIZE_DEFAULT,
|
||||
min=VAE_TILE_SIZE_MIN,
|
||||
max=VAE_TILE_SIZE_MAX,
|
||||
step=VAE_DECODE_TILE_SIZE_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_TILE_SIZE,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"overlap",
|
||||
default=VAE_OVERLAP_DEFAULT,
|
||||
min=VAE_OVERLAP_MIN,
|
||||
max=VAE_OVERLAP_MAX,
|
||||
step=VAE_OVERLAP_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_OVERLAP,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"temporal_size",
|
||||
default=VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
min=VAE_TEMPORAL_SIZE_MIN,
|
||||
max=VAE_TEMPORAL_SIZE_MAX,
|
||||
step=VAE_TEMPORAL_SIZE_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_DECODE_TEMPORAL_SIZE,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"temporal_overlap",
|
||||
default=VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
min=VAE_TEMPORAL_OVERLAP_MIN,
|
||||
max=VAE_TEMPORAL_OVERLAP_MAX,
|
||||
step=VAE_TEMPORAL_OVERLAP_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_TEMPORAL_OVERLAP,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Image.Output(
|
||||
"image",
|
||||
tooltip=tooltips.VAE_OPTIONS_IMAGE_OUTPUT,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
use_tiling: bool,
|
||||
samples: object,
|
||||
vae: object,
|
||||
tile_size: int,
|
||||
overlap: int,
|
||||
temporal_size: int,
|
||||
temporal_overlap: int,
|
||||
) -> Any:
|
||||
"""Expand through the legacy VAE Decode (Options) implementation."""
|
||||
|
||||
return VAEDecodeOptions().decode(
|
||||
use_tiling=use_tiling,
|
||||
samples=samples,
|
||||
vae=vae,
|
||||
tile_size=tile_size,
|
||||
overlap=overlap,
|
||||
temporal_size=temporal_size,
|
||||
temporal_overlap=temporal_overlap,
|
||||
)
|
||||
@@ -0,0 +1,143 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Comfy v3 node wrapper for VAE Encode (Options)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..nodes import tooltips
|
||||
from ..nodes.vae_options import (
|
||||
VAE_ENCODE_TILE_SIZE_STEP,
|
||||
VAE_OVERLAP_DEFAULT,
|
||||
VAE_OVERLAP_MAX,
|
||||
VAE_OVERLAP_MIN,
|
||||
VAE_OVERLAP_STEP,
|
||||
VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
VAE_TEMPORAL_OVERLAP_MAX,
|
||||
VAE_TEMPORAL_OVERLAP_MIN,
|
||||
VAE_TEMPORAL_OVERLAP_STEP,
|
||||
VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
VAE_TEMPORAL_SIZE_MAX,
|
||||
VAE_TEMPORAL_SIZE_MIN,
|
||||
VAE_TEMPORAL_SIZE_STEP,
|
||||
VAE_TILE_SIZE_DEFAULT,
|
||||
VAE_TILE_SIZE_MAX,
|
||||
VAE_TILE_SIZE_MIN,
|
||||
VAEEncodeOptions,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
||||
class _ComfyNodeBase:
|
||||
"""Type-checking base for Comfy v3 nodes."""
|
||||
|
||||
RETURN_TYPES: ClassVar[list[str]]
|
||||
RETURN_NAMES: ClassVar[list[str]]
|
||||
|
||||
else:
|
||||
_ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode
|
||||
|
||||
_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io
|
||||
|
||||
|
||||
class VAEEncodeOptionsV3(_ComfyNodeBase):
|
||||
"""Expose VAE Encode (Options) through Comfy's v3 extension API."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Declare the VAE Encode (Options) v3 schema."""
|
||||
|
||||
return _comfy_io.Schema(
|
||||
node_id="SimpleSyrup.VAEEncodeOptions",
|
||||
display_name="VAE Encode (Options)",
|
||||
enable_expand=True,
|
||||
category="SimpleSyrup/Latent",
|
||||
description=VAEEncodeOptions.DESCRIPTION,
|
||||
search_aliases=VAEEncodeOptions.SEARCH_ALIASES,
|
||||
inputs=[
|
||||
_comfy_io.Boolean.Input(
|
||||
"use_tiling",
|
||||
default=False,
|
||||
tooltip=tooltips.VAE_OPTIONS_USE_TILING,
|
||||
),
|
||||
_comfy_io.Image.Input(
|
||||
"pixels",
|
||||
raw_link=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_ENCODE_PIXELS,
|
||||
),
|
||||
_comfy_io.Vae.Input(
|
||||
"vae",
|
||||
raw_link=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_VAE,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"tile_size",
|
||||
default=VAE_TILE_SIZE_DEFAULT,
|
||||
min=VAE_TILE_SIZE_MIN,
|
||||
max=VAE_TILE_SIZE_MAX,
|
||||
step=VAE_ENCODE_TILE_SIZE_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_TILE_SIZE,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"overlap",
|
||||
default=VAE_OVERLAP_DEFAULT,
|
||||
min=VAE_OVERLAP_MIN,
|
||||
max=VAE_OVERLAP_MAX,
|
||||
step=VAE_OVERLAP_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_OVERLAP,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"temporal_size",
|
||||
default=VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
min=VAE_TEMPORAL_SIZE_MIN,
|
||||
max=VAE_TEMPORAL_SIZE_MAX,
|
||||
step=VAE_TEMPORAL_SIZE_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_ENCODE_TEMPORAL_SIZE,
|
||||
),
|
||||
_comfy_io.Int.Input(
|
||||
"temporal_overlap",
|
||||
default=VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
min=VAE_TEMPORAL_OVERLAP_MIN,
|
||||
max=VAE_TEMPORAL_OVERLAP_MAX,
|
||||
step=VAE_TEMPORAL_OVERLAP_STEP,
|
||||
advanced=True,
|
||||
tooltip=tooltips.VAE_OPTIONS_TEMPORAL_OVERLAP,
|
||||
),
|
||||
],
|
||||
outputs=[
|
||||
_comfy_io.Latent.Output(
|
||||
"latent",
|
||||
tooltip=tooltips.VAE_OPTIONS_LATENT_OUTPUT,
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(
|
||||
cls,
|
||||
use_tiling: bool,
|
||||
pixels: object,
|
||||
vae: object,
|
||||
tile_size: int,
|
||||
overlap: int,
|
||||
temporal_size: int,
|
||||
temporal_overlap: int,
|
||||
) -> Any:
|
||||
"""Expand through the legacy VAE Encode (Options) implementation."""
|
||||
|
||||
return VAEEncodeOptions().encode(
|
||||
use_tiling=use_tiling,
|
||||
pixels=pixels,
|
||||
vae=vae,
|
||||
tile_size=tile_size,
|
||||
overlap=overlap,
|
||||
temporal_size=temporal_size,
|
||||
temporal_overlap=temporal_overlap,
|
||||
)
|
||||
@@ -18,6 +18,12 @@ from ..domain.graph_provenance import (
|
||||
|
||||
MAX_PROVENANCE_HOPS = 128
|
||||
PASSTHROUGH_ATTRIBUTE = "GRAPH_PASSTHROUGH_OUTPUTS"
|
||||
VAE_DECODE_SOURCE_CLASS_TYPES = frozenset(
|
||||
{
|
||||
"VAEDecode",
|
||||
"SimpleSyrup.VAEDecodeOptions",
|
||||
}
|
||||
)
|
||||
|
||||
PromptNode = Mapping[str, Any]
|
||||
PromptGraph = Mapping[str, Any]
|
||||
@@ -67,8 +73,8 @@ def trace_vae_decode_provenance(
|
||||
class_type=class_type,
|
||||
)
|
||||
|
||||
if class_type == "VAEDecode":
|
||||
return _trace_vae_decode(node_id, output_slot, inputs)
|
||||
if class_type in VAE_DECODE_SOURCE_CLASS_TYPES:
|
||||
return _trace_vae_decode(node_id, output_slot, inputs, class_type)
|
||||
|
||||
class_def = node_registry.get(class_type)
|
||||
if class_def is None:
|
||||
@@ -135,15 +141,16 @@ def _trace_vae_decode(
|
||||
node_id: str,
|
||||
output_slot: int,
|
||||
inputs: Mapping[str, Any],
|
||||
class_type: str,
|
||||
) -> ProvenanceTrace:
|
||||
"""Resolve the latent and VAE links from a `VAEDecode` prompt node."""
|
||||
"""Resolve latent and VAE links from a trusted VAE decode prompt node."""
|
||||
|
||||
image_output = (node_id, output_slot)
|
||||
if output_slot != 0:
|
||||
return BrokenProvenance(
|
||||
"VAEDecode output is not the image output",
|
||||
node_id=node_id,
|
||||
class_type="VAEDecode",
|
||||
class_type=class_type,
|
||||
)
|
||||
|
||||
samples_link = parse_graph_link(inputs.get("samples"))
|
||||
@@ -151,7 +158,7 @@ def _trace_vae_decode(
|
||||
return BrokenProvenance(
|
||||
"VAEDecode samples input is not a graph link",
|
||||
node_id=node_id,
|
||||
class_type="VAEDecode",
|
||||
class_type=class_type,
|
||||
)
|
||||
|
||||
return VaeDecodeProvenance(
|
||||
|
||||
@@ -13,6 +13,7 @@ import torch
|
||||
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
|
||||
from .differential_diffusion import clone_with_differential_diffusion
|
||||
|
||||
Latent: TypeAlias = dict[str, Any]
|
||||
|
||||
@@ -108,21 +109,7 @@ class DetailSampler:
|
||||
def apply_differential_diffusion(self, model: Any) -> Any:
|
||||
"""Patch a model for feathered denoise masks when ComfyUI supports it."""
|
||||
|
||||
options = getattr(model, "model_options", {})
|
||||
if (
|
||||
isinstance(options, dict)
|
||||
and options.get("denoise_mask_function") is not None
|
||||
):
|
||||
return model
|
||||
|
||||
module = import_module("comfy_extras.nodes_differential_diffusion")
|
||||
node = module.DifferentialDiffusion
|
||||
output = node.execute(model, 1.0)
|
||||
if hasattr(output, "result"):
|
||||
return output.result[0]
|
||||
if isinstance(output, tuple):
|
||||
return output[0]
|
||||
return output[0]
|
||||
return clone_with_differential_diffusion(model)
|
||||
|
||||
|
||||
def _nodes() -> Any:
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Adapters for ComfyUI differential denoise-mask behavior."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
|
||||
def has_denoise_mask_function(model: Any) -> bool:
|
||||
"""Return whether a model patcher already has denoise-mask behavior."""
|
||||
|
||||
options = getattr(model, "model_options", {})
|
||||
return (
|
||||
isinstance(options, dict) and options.get("denoise_mask_function") is not None
|
||||
)
|
||||
|
||||
|
||||
def clone_with_differential_diffusion(model: Any, strength: float = 1.0) -> Any:
|
||||
"""Return a clone patched with ComfyUI differential denoise masks."""
|
||||
|
||||
if has_denoise_mask_function(model):
|
||||
return model
|
||||
cloned_model = model.clone()
|
||||
install_differential_diffusion(cloned_model, strength=strength)
|
||||
return cloned_model
|
||||
|
||||
|
||||
def install_differential_diffusion(model: Any, strength: float = 1.0) -> Any:
|
||||
"""Install ComfyUI differential denoise-mask behavior on a model patcher."""
|
||||
|
||||
if has_denoise_mask_function(model):
|
||||
return model
|
||||
set_mask_function = getattr(model, "set_model_denoise_mask_function", None)
|
||||
if not callable(set_mask_function):
|
||||
raise ValueError("Model does not support differential diffusion denoise masks.")
|
||||
differential_diffusion = import_module(
|
||||
"comfy_extras.nodes_differential_diffusion"
|
||||
).DifferentialDiffusion
|
||||
set_mask_function(
|
||||
lambda *args, **kwargs: differential_diffusion.forward(
|
||||
*args,
|
||||
**kwargs,
|
||||
strength=strength,
|
||||
)
|
||||
)
|
||||
return model
|
||||
@@ -0,0 +1,241 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""OpenAI-compatible HTTP client for external LLM providers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
from ..domain.external_llm import (
|
||||
ExternalLLMChatRequest,
|
||||
ExternalLLMChatResponse,
|
||||
ExternalLLMProviderError,
|
||||
normalize_base_url,
|
||||
)
|
||||
from ..shared.logging import get_logger
|
||||
|
||||
LOGGER = get_logger(__name__)
|
||||
MODELS_TIMEOUT_SECONDS = 10
|
||||
CHAT_TIMEOUT_SECONDS = 60
|
||||
|
||||
|
||||
class UrlopenResponse(Protocol):
|
||||
"""Subset of urllib response behavior used by the client."""
|
||||
|
||||
status: int
|
||||
|
||||
def read(self) -> bytes:
|
||||
"""Read response bytes."""
|
||||
|
||||
def __enter__(self) -> UrlopenResponse:
|
||||
"""Enter a response context manager."""
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc: BaseException | None,
|
||||
traceback: object | None,
|
||||
) -> None:
|
||||
"""Exit a response context manager."""
|
||||
|
||||
|
||||
Urlopen = Callable[[urllib.request.Request, int], UrlopenResponse]
|
||||
|
||||
|
||||
class ExternalLLMClient:
|
||||
"""Call OpenAI-compatible model listing and chat completion endpoints."""
|
||||
|
||||
def __init__(self, urlopen: Urlopen | None = None) -> None:
|
||||
"""Create a client with injectable HTTP transport."""
|
||||
|
||||
self._urlopen = urlopen or _default_urlopen
|
||||
|
||||
def list_models(self, base_url: str, api_key: str) -> tuple[str, ...]:
|
||||
"""Return provider model ids from `/models`."""
|
||||
|
||||
payload = self._request_json(
|
||||
base_url=base_url,
|
||||
path="/models",
|
||||
api_key=api_key,
|
||||
method="GET",
|
||||
body=None,
|
||||
timeout=MODELS_TIMEOUT_SECONDS,
|
||||
)
|
||||
return parse_models_payload(payload)
|
||||
|
||||
def create_chat_completion(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request: ExternalLLMChatRequest,
|
||||
) -> ExternalLLMChatResponse:
|
||||
"""Return assistant content from `/chat/completions`."""
|
||||
|
||||
messages: list[dict[str, object]] = [
|
||||
{"role": "system", "content": request.system_prompt},
|
||||
{"role": "user", "content": _user_message_content(request)},
|
||||
]
|
||||
payload = self._request_json(
|
||||
base_url=base_url,
|
||||
path="/chat/completions",
|
||||
api_key=api_key,
|
||||
method="POST",
|
||||
body=_chat_completion_body(request, messages),
|
||||
timeout=CHAT_TIMEOUT_SECONDS,
|
||||
)
|
||||
return ExternalLLMChatResponse(parse_chat_content(payload))
|
||||
|
||||
def _request_json(
|
||||
self,
|
||||
base_url: str,
|
||||
path: str,
|
||||
api_key: str,
|
||||
method: str,
|
||||
body: dict[str, object] | None,
|
||||
timeout: int,
|
||||
) -> object:
|
||||
"""Send an authenticated provider request and decode JSON."""
|
||||
|
||||
endpoint = f"{normalize_base_url(base_url)}{path}"
|
||||
data = None if body is None else json.dumps(body).encode("utf-8")
|
||||
request = urllib.request.Request(
|
||||
endpoint,
|
||||
data=data,
|
||||
method=method,
|
||||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
with self._urlopen(request, timeout) as response:
|
||||
raw = response.read()
|
||||
except urllib.error.HTTPError as error:
|
||||
LOGGER.warning(
|
||||
"external llm provider returned http error",
|
||||
extra={"operation": path, "status": error.code},
|
||||
)
|
||||
raise ExternalLLMProviderError(
|
||||
f"External LLM provider returned HTTP {error.code}."
|
||||
) from error
|
||||
except (urllib.error.URLError, TimeoutError) as error:
|
||||
LOGGER.warning(
|
||||
"external llm provider request failed",
|
||||
extra={"operation": path, "reason": str(error)},
|
||||
)
|
||||
raise ExternalLLMProviderError(
|
||||
"External LLM provider request failed before a response was received."
|
||||
) from error
|
||||
|
||||
try:
|
||||
return json.loads(raw.decode("utf-8"))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError) as error:
|
||||
raise ExternalLLMProviderError(
|
||||
"External LLM provider returned invalid JSON."
|
||||
) from error
|
||||
|
||||
|
||||
def parse_models_payload(payload: object) -> tuple[str, ...]:
|
||||
"""Extract de-duplicated model ids from an OpenAI-compatible payload."""
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise ExternalLLMProviderError(
|
||||
"External LLM provider returned an invalid models response."
|
||||
)
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, list):
|
||||
raise ExternalLLMProviderError(
|
||||
"External LLM provider returned an invalid models response."
|
||||
)
|
||||
|
||||
models: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for item in data:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
model_id: Any = item.get("id")
|
||||
if not isinstance(model_id, str):
|
||||
continue
|
||||
model = model_id.strip()
|
||||
if model and model not in seen:
|
||||
seen.add(model)
|
||||
models.append(model)
|
||||
return tuple(models)
|
||||
|
||||
|
||||
def _default_urlopen(
|
||||
request: urllib.request.Request,
|
||||
timeout: int,
|
||||
) -> UrlopenResponse:
|
||||
"""Open a URL request with an explicit timeout."""
|
||||
|
||||
return cast(UrlopenResponse, urllib.request.urlopen(request, timeout=timeout))
|
||||
|
||||
|
||||
def _user_message_content(request: ExternalLLMChatRequest) -> object:
|
||||
"""Return text-only or multimodal user message content."""
|
||||
|
||||
if request.image_data_url is None:
|
||||
return request.user_prompt
|
||||
return [
|
||||
{"type": "text", "text": request.user_prompt},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": request.image_data_url},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _chat_completion_body(
|
||||
request: ExternalLLMChatRequest,
|
||||
messages: list[dict[str, object]],
|
||||
) -> dict[str, object]:
|
||||
"""Return an OpenAI-compatible chat request body with optional reasoning."""
|
||||
|
||||
body: dict[str, object] = {
|
||||
"model": request.model,
|
||||
"messages": messages,
|
||||
"max_tokens": request.max_tokens,
|
||||
}
|
||||
if request.reasoning_effort in {"high", "medium", "low"}:
|
||||
body["reasoning_effort"] = request.reasoning_effort
|
||||
if request.reasoning_effort == "off":
|
||||
body["chat_template_kwargs"] = {"thinking": False}
|
||||
return body
|
||||
|
||||
|
||||
def parse_chat_content(payload: object) -> str:
|
||||
"""Extract assistant message content from a chat completion payload."""
|
||||
|
||||
if not isinstance(payload, dict):
|
||||
raise ExternalLLMProviderError(
|
||||
"The external LLM provider returned an invalid chat completion response."
|
||||
)
|
||||
choices = payload.get("choices")
|
||||
if not isinstance(choices, list) or not choices:
|
||||
raise ExternalLLMProviderError(
|
||||
"The external LLM provider returned an invalid chat completion response."
|
||||
)
|
||||
first = choices[0]
|
||||
if not isinstance(first, dict):
|
||||
raise ExternalLLMProviderError(
|
||||
"The external LLM provider returned an invalid chat completion response."
|
||||
)
|
||||
message = first.get("message")
|
||||
if not isinstance(message, dict):
|
||||
raise ExternalLLMProviderError(
|
||||
"The external LLM provider returned an invalid chat completion response."
|
||||
)
|
||||
content = message.get("content")
|
||||
if not isinstance(content, str):
|
||||
raise ExternalLLMProviderError(
|
||||
"The external LLM provider returned an invalid chat completion response."
|
||||
)
|
||||
return content
|
||||
@@ -0,0 +1,150 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Image adaptation for OpenAI-compatible external LLM vision requests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from io import BytesIO
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from ..domain.segs import Segment
|
||||
from ..masking.segs_mask_ops import crop_image, crop_mask, resize_mask
|
||||
from ..shared.tensor_validation import validate_image_tensor
|
||||
|
||||
SEG_IMAGE_MODES = ("transparent mask", "black mask", "full crop")
|
||||
|
||||
|
||||
class ExternalLLMImageEncoder:
|
||||
"""Encode ComfyUI IMAGE tensors for OpenAI-compatible vision payloads."""
|
||||
|
||||
def encode_first_image_as_data_url(self, image: object) -> str:
|
||||
"""Return the first image in a ComfyUI IMAGE batch as a PNG data URL."""
|
||||
|
||||
validate_image_tensor(image)
|
||||
if not isinstance(image, torch.Tensor):
|
||||
raise TypeError("image must be a torch.Tensor with shape (B, H, W, C).")
|
||||
|
||||
pil_image = _image_tensor_to_rgb_pil(image)
|
||||
buffer = BytesIO()
|
||||
pil_image.save(buffer, format="PNG")
|
||||
encoded = base64.b64encode(buffer.getvalue()).decode("ascii")
|
||||
return f"data:image/png;base64,{encoded}"
|
||||
|
||||
|
||||
class ExternalLLMSegsImageEncoder:
|
||||
"""Encode SEG crops for OpenAI-compatible vision payloads."""
|
||||
|
||||
def encode_segment_as_data_url(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
segment: Segment,
|
||||
mode: str,
|
||||
) -> str:
|
||||
"""Return one SEG crop as a PNG data URL."""
|
||||
|
||||
if mode not in SEG_IMAGE_MODES:
|
||||
choices = ", ".join(SEG_IMAGE_MODES)
|
||||
raise ValueError(f"seg_image_mode must be one of: {choices}.")
|
||||
validate_image_tensor(image)
|
||||
if int(image.shape[0]) != 1:
|
||||
raise ValueError("SEG crop image encoding requires one IMAGE item.")
|
||||
|
||||
crop = crop_image(image.float().clamp(0.0, 1.0), segment.crop_region)
|
||||
if mode == "full crop":
|
||||
return _pil_to_png_data_url(_image_tensor_to_rgb_pil(crop))
|
||||
|
||||
mask = _segment_mask_for_crop(
|
||||
segment=segment,
|
||||
image_height=int(image.shape[1]),
|
||||
image_width=int(image.shape[2]),
|
||||
)
|
||||
if int(mask.shape[0]) != int(crop.shape[1]) or int(mask.shape[1]) != int(
|
||||
crop.shape[2]
|
||||
):
|
||||
mask = resize_mask(mask, int(crop.shape[1]), int(crop.shape[2]))
|
||||
|
||||
if mode == "black mask":
|
||||
masked_crop = crop * mask.unsqueeze(0).unsqueeze(-1)
|
||||
return _pil_to_png_data_url(_image_tensor_to_rgb_pil(masked_crop))
|
||||
|
||||
return _pil_to_png_data_url(_crop_and_mask_to_rgba_pil(crop, mask))
|
||||
|
||||
|
||||
def _image_tensor_to_rgb_pil(image: torch.Tensor) -> Image.Image:
|
||||
"""Convert the first BHWC image tensor item to RGB PIL image."""
|
||||
|
||||
import numpy as np
|
||||
|
||||
array = image[0].detach().cpu().float().clamp(0.0, 1.0).numpy()
|
||||
if array.shape[-1] == 1:
|
||||
array = np.repeat(array, 3, axis=-1)
|
||||
elif array.shape[-1] >= 3:
|
||||
array = array[..., :3]
|
||||
else:
|
||||
array = np.repeat(array[..., :1], 3, axis=-1)
|
||||
return Image.fromarray((array * 255.0).round().astype(np.uint8))
|
||||
|
||||
|
||||
def _segment_mask_for_crop(
|
||||
segment: Segment,
|
||||
image_height: int,
|
||||
image_width: int,
|
||||
) -> torch.Tensor:
|
||||
"""Return the segment mask normalized to the crop region."""
|
||||
|
||||
if not isinstance(segment.cropped_mask, torch.Tensor):
|
||||
raise TypeError("SEG cropped_mask must be a torch.Tensor.")
|
||||
|
||||
mask = _normalize_mask_shape(segment.cropped_mask)
|
||||
region = segment.crop_region
|
||||
crop_height = region.height
|
||||
crop_width = region.width
|
||||
mask_height = int(mask.shape[0])
|
||||
mask_width = int(mask.shape[1])
|
||||
if mask_height == crop_height and mask_width == crop_width:
|
||||
return mask.float().clamp(0.0, 1.0)
|
||||
if mask_height == image_height and mask_width == image_width:
|
||||
return crop_mask(mask, region).float().clamp(0.0, 1.0)
|
||||
return resize_mask(mask, crop_height, crop_width).float().clamp(0.0, 1.0)
|
||||
|
||||
|
||||
def _normalize_mask_shape(mask: torch.Tensor) -> torch.Tensor:
|
||||
"""Return a SEG mask as an HW tensor."""
|
||||
|
||||
working = mask.detach().cpu().float()
|
||||
if working.ndim == 2:
|
||||
return working
|
||||
if working.ndim == 3:
|
||||
return working[0]
|
||||
raise ValueError("SEG cropped_mask must be an HW or BHW tensor.")
|
||||
|
||||
|
||||
def _crop_and_mask_to_rgba_pil(crop: torch.Tensor, mask: torch.Tensor) -> Image.Image:
|
||||
"""Convert a BHWC crop and HW alpha mask to an RGBA PIL image."""
|
||||
|
||||
import numpy as np
|
||||
|
||||
rgb = crop[0].detach().cpu().float().clamp(0.0, 1.0).numpy()
|
||||
if rgb.shape[-1] == 1:
|
||||
rgb = np.repeat(rgb, 3, axis=-1)
|
||||
elif rgb.shape[-1] >= 3:
|
||||
rgb = rgb[..., :3]
|
||||
else:
|
||||
rgb = np.repeat(rgb[..., :1], 3, axis=-1)
|
||||
alpha = mask.detach().cpu().float().clamp(0.0, 1.0).numpy()
|
||||
rgba = np.concatenate((rgb, alpha[..., None]), axis=-1)
|
||||
return Image.fromarray((rgba * 255.0).round().astype(np.uint8), mode="RGBA")
|
||||
|
||||
|
||||
def _pil_to_png_data_url(image: Image.Image) -> str:
|
||||
"""Return a PNG data URL for a PIL image."""
|
||||
|
||||
buffer = BytesIO()
|
||||
image.save(buffer, format="PNG")
|
||||
encoded = base64.b64encode(buffer.getvalue()).decode("ascii")
|
||||
return f"data:image/png;base64,{encoded}"
|
||||
@@ -0,0 +1,102 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""OS credential storage adapter for external LLM API keys."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from typing import Protocol, cast
|
||||
|
||||
from ..domain.external_llm import ExternalLLMConfigError, normalize_base_url
|
||||
|
||||
SERVICE_NAME = "SimpleSyrup"
|
||||
|
||||
|
||||
class KeyringProtocol(Protocol):
|
||||
"""Subset of the keyring API used by SimpleSyrup."""
|
||||
|
||||
def get_password(self, service_name: str, username: str) -> str | None:
|
||||
"""Return a stored password or None."""
|
||||
|
||||
def set_password(self, service_name: str, username: str, password: str) -> None:
|
||||
"""Store a password."""
|
||||
|
||||
def delete_password(self, service_name: str, username: str) -> None:
|
||||
"""Delete a stored password."""
|
||||
|
||||
|
||||
class ExternalLLMKeyringError(RuntimeError):
|
||||
"""Raised when OS credential storage cannot complete an operation."""
|
||||
|
||||
|
||||
class ExternalLLMKeyringStore:
|
||||
"""Store external LLM API keys in the OS credential backend."""
|
||||
|
||||
def __init__(self, keyring_module: KeyringProtocol | None = None) -> None:
|
||||
"""Create a keyring store with an injectable credential backend."""
|
||||
|
||||
self._keyring = keyring_module
|
||||
|
||||
def has_api_key(self, base_url: str) -> bool:
|
||||
"""Return whether an API key exists for the normalized endpoint."""
|
||||
|
||||
return self.get_api_key(base_url) != ""
|
||||
|
||||
def get_api_key(self, base_url: str) -> str:
|
||||
"""Return the API key for the endpoint or an empty string."""
|
||||
|
||||
username = credential_username(base_url)
|
||||
try:
|
||||
return self._backend().get_password(SERVICE_NAME, username) or ""
|
||||
except Exception as error:
|
||||
raise ExternalLLMKeyringError(
|
||||
"Could not read the external LLM API key from OS credential storage."
|
||||
) from error
|
||||
|
||||
def save_api_key(self, base_url: str, api_key: str) -> None:
|
||||
"""Store an API key for the normalized endpoint."""
|
||||
|
||||
secret = api_key.strip()
|
||||
if not secret:
|
||||
raise ExternalLLMConfigError("External LLM API key must not be empty.")
|
||||
|
||||
username = credential_username(base_url)
|
||||
try:
|
||||
self._backend().set_password(SERVICE_NAME, username, secret)
|
||||
except Exception as error:
|
||||
raise ExternalLLMKeyringError(
|
||||
"Could not save the external LLM API key in OS credential storage."
|
||||
) from error
|
||||
|
||||
def delete_api_key(self, base_url: str) -> None:
|
||||
"""Delete the API key for the normalized endpoint when present."""
|
||||
|
||||
username = credential_username(base_url)
|
||||
try:
|
||||
self._backend().delete_password(SERVICE_NAME, username)
|
||||
except Exception as error:
|
||||
raise ExternalLLMKeyringError(
|
||||
"Could not delete the external LLM API key from OS credential storage."
|
||||
) from error
|
||||
|
||||
def _backend(self) -> KeyringProtocol:
|
||||
"""Return the configured keyring backend."""
|
||||
|
||||
if self._keyring is not None:
|
||||
return self._keyring
|
||||
|
||||
try:
|
||||
module = importlib.import_module("keyring")
|
||||
except ModuleNotFoundError as error:
|
||||
raise ExternalLLMKeyringError(
|
||||
"The keyring package is required to store external LLM API keys."
|
||||
) from error
|
||||
return cast(KeyringProtocol, module)
|
||||
|
||||
|
||||
def credential_username(base_url: str) -> str:
|
||||
"""Return the stable keyring username for an endpoint."""
|
||||
|
||||
return f"external-llm:{normalize_base_url(base_url)}"
|
||||
@@ -0,0 +1,196 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""HTTP route registration for external LLM provider settings."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
from collections.abc import Callable, Coroutine
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from ..domain.external_llm import ExternalLLMConfigError, ExternalLLMProviderError
|
||||
from ..services.external_llm_prompt_service import ExternalLLMPromptService
|
||||
from ..shared.logging import get_logger
|
||||
from .external_llm_keyring import ExternalLLMKeyringError
|
||||
|
||||
LOGGER = get_logger(__name__)
|
||||
EXTERNAL_LLM_SETTINGS_ROUTE = "/simple-syrup/external-llm/settings"
|
||||
EXTERNAL_LLM_API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"
|
||||
EXTERNAL_LLM_MODELS_REFRESH_ROUTE = "/simple-syrup/external-llm/models/refresh"
|
||||
|
||||
Handler = Callable[[Any], Coroutine[Any, Any, web.Response]]
|
||||
_REGISTERED_PROMPT_SERVERS: set[int] = set()
|
||||
|
||||
|
||||
class RoutesProtocol(Protocol):
|
||||
"""Subset of Comfy's route table needed for route registration."""
|
||||
|
||||
def get(self, path: str) -> Callable[[Handler], Handler]:
|
||||
"""Return a GET route decorator."""
|
||||
|
||||
def post(self, path: str) -> Callable[[Handler], Handler]:
|
||||
"""Return a POST route decorator."""
|
||||
|
||||
def delete(self, path: str) -> Callable[[Handler], Handler]:
|
||||
"""Return a DELETE route decorator."""
|
||||
|
||||
|
||||
class PromptServerProtocol(Protocol):
|
||||
"""Subset of Comfy's PromptServer needed for route registration."""
|
||||
|
||||
routes: RoutesProtocol
|
||||
|
||||
|
||||
class ExternalLLMHandlers:
|
||||
"""Handle SimpleSyrup external LLM HTTP requests."""
|
||||
|
||||
def __init__(self, service: ExternalLLMPromptService) -> None:
|
||||
"""Create handlers backed by the external LLM service."""
|
||||
|
||||
self._service = service
|
||||
|
||||
async def get_settings(self, _request: Any) -> web.Response:
|
||||
"""Return current non-secret external LLM settings."""
|
||||
|
||||
return web.json_response(self._settings_payload())
|
||||
|
||||
async def post_settings(self, request: Any) -> web.Response:
|
||||
"""Validate and persist non-secret external LLM settings."""
|
||||
|
||||
try:
|
||||
payload = await request.json()
|
||||
if not isinstance(payload, dict):
|
||||
raise ExternalLLMConfigError(
|
||||
"External LLM settings request body must be a JSON object."
|
||||
)
|
||||
base_url = payload.get("base_url", "")
|
||||
default_model = payload.get("default_model", "")
|
||||
if not isinstance(base_url, str) or not isinstance(default_model, str):
|
||||
raise ExternalLLMConfigError(
|
||||
"External LLM settings require string base_url and "
|
||||
"default_model values."
|
||||
)
|
||||
self._service.save_config(base_url, default_model)
|
||||
except ExternalLLMConfigError as error:
|
||||
return web.json_response({"error": str(error)}, status=400)
|
||||
except Exception as error:
|
||||
LOGGER.warning(
|
||||
"invalid external llm settings request body",
|
||||
extra={"route": EXTERNAL_LLM_SETTINGS_ROUTE, "reason": str(error)},
|
||||
)
|
||||
return web.json_response(
|
||||
{"error": "External LLM settings request body must be valid JSON."},
|
||||
status=400,
|
||||
)
|
||||
|
||||
return web.json_response(self._settings_payload())
|
||||
|
||||
async def post_api_key(self, request: Any) -> web.Response:
|
||||
"""Store an external LLM API key without returning it."""
|
||||
|
||||
try:
|
||||
payload = await request.json()
|
||||
if not isinstance(payload, dict):
|
||||
raise ExternalLLMConfigError(
|
||||
"External LLM API key request body must be a JSON object."
|
||||
)
|
||||
api_key = payload.get("api_key")
|
||||
if not isinstance(api_key, str):
|
||||
raise ExternalLLMConfigError(
|
||||
"External LLM API key request requires an api_key string."
|
||||
)
|
||||
await asyncio.to_thread(self._service.save_api_key, api_key)
|
||||
except (
|
||||
ExternalLLMConfigError,
|
||||
ExternalLLMKeyringError,
|
||||
ExternalLLMProviderError,
|
||||
) as error:
|
||||
return web.json_response({"error": str(error)}, status=400)
|
||||
except Exception as error:
|
||||
LOGGER.warning(
|
||||
"invalid external llm api key request body",
|
||||
extra={"route": EXTERNAL_LLM_API_KEY_ROUTE, "reason": str(error)},
|
||||
)
|
||||
return web.json_response(
|
||||
{"error": "External LLM API key request body must be valid JSON."},
|
||||
status=400,
|
||||
)
|
||||
|
||||
return web.json_response(self._settings_payload())
|
||||
|
||||
async def delete_api_key(self, _request: Any) -> web.Response:
|
||||
"""Delete the configured external LLM API key."""
|
||||
|
||||
try:
|
||||
self._service.delete_api_key()
|
||||
except ExternalLLMKeyringError as error:
|
||||
return web.json_response({"error": str(error)}, status=400)
|
||||
|
||||
return web.json_response(self._settings_payload())
|
||||
|
||||
async def refresh_models(self, _request: Any) -> web.Response:
|
||||
"""Refresh cached external LLM model ids."""
|
||||
|
||||
try:
|
||||
await asyncio.to_thread(self._service.refresh_models)
|
||||
except (
|
||||
ExternalLLMConfigError,
|
||||
ExternalLLMKeyringError,
|
||||
ExternalLLMProviderError,
|
||||
) as error:
|
||||
return web.json_response({"error": str(error)}, status=400)
|
||||
|
||||
return web.json_response(self._settings_payload())
|
||||
|
||||
def _settings_payload(self) -> dict[str, object]:
|
||||
"""Return non-secret external LLM settings plus API key presence."""
|
||||
|
||||
return self._service.settings_payload()
|
||||
|
||||
|
||||
def register_external_llm_routes(
|
||||
service: ExternalLLMPromptService | None = None,
|
||||
prompt_server: PromptServerProtocol | None = None,
|
||||
) -> bool:
|
||||
"""Register external LLM routes with Comfy's PromptServer."""
|
||||
|
||||
server_instance = prompt_server or _prompt_server_instance()
|
||||
if server_instance is None:
|
||||
return False
|
||||
|
||||
server_key = id(server_instance)
|
||||
if prompt_server is None and server_key in _REGISTERED_PROMPT_SERVERS:
|
||||
return True
|
||||
|
||||
handlers = ExternalLLMHandlers(service or ExternalLLMPromptService())
|
||||
server_instance.routes.get(EXTERNAL_LLM_SETTINGS_ROUTE)(handlers.get_settings)
|
||||
server_instance.routes.post(EXTERNAL_LLM_SETTINGS_ROUTE)(handlers.post_settings)
|
||||
server_instance.routes.post(EXTERNAL_LLM_API_KEY_ROUTE)(handlers.post_api_key)
|
||||
server_instance.routes.delete(EXTERNAL_LLM_API_KEY_ROUTE)(handlers.delete_api_key)
|
||||
server_instance.routes.post(EXTERNAL_LLM_MODELS_REFRESH_ROUTE)(
|
||||
handlers.refresh_models
|
||||
)
|
||||
if prompt_server is None:
|
||||
_REGISTERED_PROMPT_SERVERS.add(server_key)
|
||||
return True
|
||||
|
||||
|
||||
def _prompt_server_instance() -> PromptServerProtocol | None:
|
||||
"""Return Comfy's PromptServer instance when available."""
|
||||
|
||||
try:
|
||||
server_module = sys.modules["server"]
|
||||
prompt_server = server_module.PromptServer
|
||||
instance = prompt_server.instance
|
||||
except (KeyError, AttributeError) as error:
|
||||
LOGGER.debug(
|
||||
"PromptServer unavailable for external LLM routes",
|
||||
extra={"reason": str(error)},
|
||||
)
|
||||
return None
|
||||
return cast(PromptServerProtocol, instance)
|
||||
@@ -26,6 +26,7 @@ from ..domain.tiled_diffusion import (
|
||||
from ..shared.logging import get_logger
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
|
||||
from .differential_diffusion import install_differential_diffusion
|
||||
from .tiled_sampling import (
|
||||
ApplyModel,
|
||||
Latent,
|
||||
@@ -60,6 +61,7 @@ def sample_mixture_of_diffusers(
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample a latent with a cloned model patched for Mixture of Diffusers."""
|
||||
|
||||
@@ -105,6 +107,7 @@ def sample_mixture_of_diffusers(
|
||||
tile_height=latent_tile_height,
|
||||
overlap=latent_tile_overlap,
|
||||
tile_batch_size=latent_tile_batch_size,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
|
||||
batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None
|
||||
@@ -160,6 +163,7 @@ def clone_model_with_mixture_of_diffusers(
|
||||
tile_height: int,
|
||||
overlap: int,
|
||||
tile_batch_size: int,
|
||||
differential_diffusion: bool = False,
|
||||
) -> tuple[Any, TiledDiffusionPlan]:
|
||||
"""Return a model clone patched with a pre-CFG Mixture wrapper."""
|
||||
|
||||
@@ -172,6 +176,8 @@ def clone_model_with_mixture_of_diffusers(
|
||||
tile_batch_size=tile_batch_size,
|
||||
)
|
||||
cloned_model = model.clone()
|
||||
if differential_diffusion:
|
||||
install_differential_diffusion(cloned_model)
|
||||
old_wrapper = cloned_model.model_options.get("model_function_wrapper")
|
||||
if old_wrapper is not None and not callable(old_wrapper):
|
||||
raise ValueError("Existing model_function_wrapper is not callable.")
|
||||
|
||||
@@ -25,6 +25,7 @@ from ..domain.tiled_diffusion import (
|
||||
from ..shared.logging import get_logger
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
|
||||
from .differential_diffusion import install_differential_diffusion
|
||||
from .tiled_sampling import (
|
||||
ApplyModel,
|
||||
Latent,
|
||||
@@ -60,6 +61,7 @@ def sample_multidiffusion(
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample a latent with a cloned model patched for MultiDiffusion."""
|
||||
|
||||
@@ -106,6 +108,7 @@ def sample_multidiffusion(
|
||||
tile_height=latent_tile_height,
|
||||
overlap=latent_tile_overlap,
|
||||
tile_batch_size=latent_tile_batch_size,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
|
||||
batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None
|
||||
@@ -162,6 +165,7 @@ def clone_model_with_multidiffusion(
|
||||
tile_height: int,
|
||||
overlap: int,
|
||||
tile_batch_size: int,
|
||||
differential_diffusion: bool = False,
|
||||
) -> tuple[Any, TiledDiffusionPlan]:
|
||||
"""Return a model clone patched with a pre-CFG MultiDiffusion wrapper."""
|
||||
|
||||
@@ -174,6 +178,8 @@ def clone_model_with_multidiffusion(
|
||||
tile_batch_size=tile_batch_size,
|
||||
)
|
||||
cloned_model = model.clone()
|
||||
if differential_diffusion:
|
||||
install_differential_diffusion(cloned_model)
|
||||
old_wrapper = cloned_model.model_options.get("model_function_wrapper")
|
||||
if old_wrapper is not None and not callable(old_wrapper):
|
||||
raise ValueError("Existing model_function_wrapper is not callable.")
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Runtime graph expansion for Prompt-Control scheduling and prompt encoding."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from importlib import import_module
|
||||
from typing import Any, cast
|
||||
|
||||
from ..domain.prompt_control_prompt import (
|
||||
PreparedPromptSide,
|
||||
apply_encode_style,
|
||||
prepare_prompt_side,
|
||||
)
|
||||
from .prompt_control_availability import find_prompt_control_install
|
||||
|
||||
PROMPT_CONTROL_MISSING_MESSAGE = (
|
||||
"Schedule & Encode Prompts requires comfyui-prompt-control. "
|
||||
"Install Prompt Control or remove this node from the workflow."
|
||||
)
|
||||
PROMPT_BATCH_SEPARATOR = "[SEP]"
|
||||
|
||||
|
||||
class PromptControlScheduleEncodeGraphBuilder:
|
||||
"""Build lazy Prompt-Control graphs for LoRA scheduling and prompt encoding."""
|
||||
|
||||
def build(
|
||||
self,
|
||||
model: Any,
|
||||
clip: Any,
|
||||
positive_prompt: str,
|
||||
negative_prompt: str,
|
||||
encode_style: str = "",
|
||||
) -> Any:
|
||||
"""Return an io.NodeOutput for scheduled model and encoded prompts."""
|
||||
|
||||
io, graph_utils, lazy_nodes = self._prompt_control_dependencies()
|
||||
positive_side = prepare_prompt_side(positive_prompt, PROMPT_BATCH_SEPARATOR)
|
||||
negative_side = prepare_prompt_side(negative_prompt, PROMPT_BATCH_SEPARATOR)
|
||||
|
||||
expand: dict[str, dict[str, Any]] = {}
|
||||
positive_lora = lazy_nodes.PCLazyLoraLoaderAdvanced.execute(
|
||||
model=model,
|
||||
clip=clip,
|
||||
text=positive_side.lora_tags,
|
||||
apply_hooks=True,
|
||||
tags="",
|
||||
start=0.0,
|
||||
end=1.0,
|
||||
num_steps=0,
|
||||
)
|
||||
self._merge_expand(expand, positive_lora.expand, "positive LoRA scheduling")
|
||||
|
||||
negative_lora = lazy_nodes.PCLazyLoraLoaderAdvanced.execute(
|
||||
model=positive_lora.args[0],
|
||||
clip=positive_lora.args[1],
|
||||
text=negative_side.lora_tags,
|
||||
apply_hooks=True,
|
||||
tags="",
|
||||
start=0.0,
|
||||
end=1.0,
|
||||
num_steps=0,
|
||||
)
|
||||
self._merge_expand(expand, negative_lora.expand, "negative LoRA scheduling")
|
||||
|
||||
scheduled_model = negative_lora.args[0]
|
||||
scheduled_clip = negative_lora.args[1]
|
||||
positive_conditioning = self._encode_side(
|
||||
side=positive_side,
|
||||
clip=scheduled_clip,
|
||||
encode_style=encode_style,
|
||||
graph_utils=graph_utils,
|
||||
lazy_text_encoder=lazy_nodes.PCLazyTextEncodeAdvanced,
|
||||
expand=expand,
|
||||
label="positive prompt encoding",
|
||||
)
|
||||
negative_conditioning = self._encode_side(
|
||||
side=negative_side,
|
||||
clip=scheduled_clip,
|
||||
encode_style=encode_style,
|
||||
graph_utils=graph_utils,
|
||||
lazy_text_encoder=lazy_nodes.PCLazyTextEncodeAdvanced,
|
||||
expand=expand,
|
||||
label="negative prompt encoding",
|
||||
)
|
||||
|
||||
return io.NodeOutput(
|
||||
scheduled_model,
|
||||
positive_conditioning,
|
||||
negative_conditioning,
|
||||
expand=expand,
|
||||
)
|
||||
|
||||
def _encode_side(
|
||||
self,
|
||||
*,
|
||||
side: PreparedPromptSide,
|
||||
clip: Any,
|
||||
encode_style: str,
|
||||
graph_utils: Any,
|
||||
lazy_text_encoder: Any,
|
||||
expand: dict[str, dict[str, Any]],
|
||||
label: str,
|
||||
) -> Any:
|
||||
"""Encode one prompt side and return conditioning or conditioning batch."""
|
||||
|
||||
conditioning_outputs: list[Any] = []
|
||||
for index, chunk in enumerate(side.chunks):
|
||||
text = apply_encode_style(encode_style, chunk.text)
|
||||
node_output = lazy_text_encoder.execute(
|
||||
clip=clip,
|
||||
text=text,
|
||||
tags="",
|
||||
start=0.0,
|
||||
end=1.0,
|
||||
num_steps=0,
|
||||
)
|
||||
self._merge_expand(
|
||||
expand,
|
||||
node_output.expand,
|
||||
f"{label} chunk {index}",
|
||||
)
|
||||
conditioning_outputs.append(node_output.args[0])
|
||||
|
||||
if len(conditioning_outputs) == 1:
|
||||
return conditioning_outputs[0]
|
||||
|
||||
pack_graph = graph_utils.GraphBuilder()
|
||||
current = pack_graph.node(
|
||||
"SimpleSyrup.ConditioningBatchStart",
|
||||
conditioning=conditioning_outputs[0],
|
||||
)
|
||||
for conditioning in conditioning_outputs[1:]:
|
||||
current = pack_graph.node(
|
||||
"SimpleSyrup.ConditioningBatchAppend",
|
||||
batch=current.out(0),
|
||||
conditioning=conditioning,
|
||||
)
|
||||
self._merge_expand(
|
||||
expand,
|
||||
cast(dict[str, dict[str, Any]], pack_graph.finalize()),
|
||||
f"{label} batch packing",
|
||||
)
|
||||
return current.out(0)
|
||||
|
||||
def _prompt_control_dependencies(self) -> tuple[Any, Any, Any]:
|
||||
"""Import Prompt-Control and Comfy graph helpers on demand."""
|
||||
|
||||
try:
|
||||
io = import_module("comfy_api.latest.io")
|
||||
except ModuleNotFoundError:
|
||||
comfy_api = import_module("comfy_api.latest")
|
||||
io = comfy_api.io
|
||||
try:
|
||||
graph_utils = import_module("comfy_execution.graph_utils")
|
||||
lazy_nodes = self._import_prompt_control_lazy_nodes()
|
||||
except ModuleNotFoundError as exc:
|
||||
raise RuntimeError(PROMPT_CONTROL_MISSING_MESSAGE) from exc
|
||||
return io, graph_utils, lazy_nodes
|
||||
|
||||
def _import_prompt_control_lazy_nodes(self) -> Any:
|
||||
"""Import Prompt-Control lazy nodes from installed or sibling paths."""
|
||||
|
||||
try:
|
||||
return import_module("prompt_control.nodes_lazy")
|
||||
except ModuleNotFoundError:
|
||||
availability = find_prompt_control_install()
|
||||
if availability.root_path is not None:
|
||||
root_path = str(availability.root_path)
|
||||
if root_path not in sys.path:
|
||||
sys.path.insert(0, root_path)
|
||||
return import_module("prompt_control.nodes_lazy")
|
||||
|
||||
def _merge_expand(
|
||||
self,
|
||||
target: dict[str, dict[str, Any]],
|
||||
source: object,
|
||||
operation: str,
|
||||
) -> None:
|
||||
"""Merge a lazy expand graph and reject duplicate generated node ids."""
|
||||
|
||||
if not source:
|
||||
return
|
||||
expand = cast(dict[str, dict[str, Any]], source)
|
||||
overlap = set(target).intersection(expand)
|
||||
if overlap:
|
||||
overlapping_ids = ", ".join(sorted(overlap))
|
||||
raise ValueError(
|
||||
f"Prompt-Control graph expansion generated duplicate node ids "
|
||||
f"during {operation}: {overlapping_ids}."
|
||||
)
|
||||
target.update(expand)
|
||||
@@ -22,6 +22,7 @@ from ..domain.regional_detailing import LatentRegion
|
||||
from ..shared.logging import get_logger
|
||||
from . import sampling_samplers, sampling_schedulers
|
||||
from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback
|
||||
from .differential_diffusion import install_differential_diffusion
|
||||
from .tiled_sampling import (
|
||||
Latent,
|
||||
reject_unsupported_conditioning,
|
||||
@@ -62,6 +63,7 @@ def sample_regional_multidiffusion(
|
||||
denoise: float,
|
||||
global_prompt_weight: float,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample a latent with regional MultiDiffusion prompt blending."""
|
||||
|
||||
@@ -109,6 +111,7 @@ def sample_regional_multidiffusion(
|
||||
latent_ndim=latent_samples.ndim,
|
||||
regions=regions,
|
||||
global_prompt_weight=global_prompt_weight,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
|
||||
batch_inds = latent_image["batch_index"] if "batch_index" in latent_image else None
|
||||
@@ -163,6 +166,7 @@ def clone_model_with_regional_multidiffusion(
|
||||
latent_ndim: int,
|
||||
regions: tuple[LatentRegion, ...],
|
||||
global_prompt_weight: float = 0.0,
|
||||
differential_diffusion: bool = False,
|
||||
) -> tuple[Any, RegionalMultiDiffusionSummary]:
|
||||
"""Return a model clone patched with regional calc-cond-batch blending."""
|
||||
|
||||
@@ -177,6 +181,8 @@ def clone_model_with_regional_multidiffusion(
|
||||
regions=regions,
|
||||
)
|
||||
cloned_model = model.clone()
|
||||
if differential_diffusion:
|
||||
install_differential_diffusion(cloned_model)
|
||||
old_wrapper = cloned_model.model_options.get("sampler_calc_cond_batch_function")
|
||||
if old_wrapper is not None and not callable(old_wrapper):
|
||||
raise ValueError("Existing sampler_calc_cond_batch_function is not callable.")
|
||||
|
||||
@@ -8,12 +8,17 @@ from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from json import JSONDecodeError
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
from typing import Any, Final
|
||||
|
||||
from ..domain.external_llm import (
|
||||
ExternalLLMConfigError,
|
||||
normalize_base_url,
|
||||
normalize_model_ids,
|
||||
)
|
||||
from ..shared.logging import get_logger
|
||||
|
||||
LOGGER = get_logger(__name__)
|
||||
@@ -24,16 +29,83 @@ class SimpleSyrupSettingsError(ValueError):
|
||||
"""Raised when SimpleSyrup settings data is malformed."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExternalLLMSettings:
|
||||
"""Non-secret external LLM provider settings persisted in Comfy's user data."""
|
||||
|
||||
base_url: str = ""
|
||||
cached_models: tuple[str, ...] = ()
|
||||
default_model: str = ""
|
||||
|
||||
def to_payload(self) -> dict[str, object]:
|
||||
"""Return the validated external LLM JSON payload shape."""
|
||||
|
||||
return {
|
||||
"base_url": self.base_url,
|
||||
"cached_models": list(self.cached_models),
|
||||
"default_model": self.default_model,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, payload: object) -> ExternalLLMSettings:
|
||||
"""Create external LLM settings from a decoded JSON payload."""
|
||||
|
||||
if payload is None:
|
||||
return cls()
|
||||
if not isinstance(payload, dict):
|
||||
raise SimpleSyrupSettingsError(
|
||||
"SimpleSyrup external_llm settings must be a JSON object."
|
||||
)
|
||||
|
||||
base_url = payload.get("base_url", "")
|
||||
if not isinstance(base_url, str):
|
||||
raise SimpleSyrupSettingsError(
|
||||
"SimpleSyrup external_llm.base_url must be a string."
|
||||
)
|
||||
|
||||
default_model = payload.get("default_model", "")
|
||||
if not isinstance(default_model, str):
|
||||
raise SimpleSyrupSettingsError(
|
||||
"SimpleSyrup external_llm.default_model must be a string."
|
||||
)
|
||||
|
||||
try:
|
||||
cached_models = normalize_model_ids(payload.get("cached_models", []))
|
||||
except ExternalLLMConfigError as error:
|
||||
raise SimpleSyrupSettingsError(str(error)) from error
|
||||
|
||||
default = default_model.strip()
|
||||
if default and cached_models and default not in cached_models:
|
||||
default = cached_models[0]
|
||||
|
||||
normalized_url = ""
|
||||
if base_url.strip():
|
||||
try:
|
||||
normalized_url = normalize_base_url(base_url)
|
||||
except ExternalLLMConfigError as error:
|
||||
raise SimpleSyrupSettingsError(str(error)) from error
|
||||
|
||||
return cls(
|
||||
base_url=normalized_url,
|
||||
cached_models=cached_models,
|
||||
default_model=default,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SimpleSyrupSettings:
|
||||
"""User-configurable SimpleSyrup runtime settings."""
|
||||
|
||||
show_downloadable_models: bool = True
|
||||
external_llm: ExternalLLMSettings = field(default_factory=ExternalLLMSettings)
|
||||
|
||||
def to_payload(self) -> dict[str, bool]:
|
||||
def to_payload(self) -> dict[str, object]:
|
||||
"""Return the validated JSON payload shape."""
|
||||
|
||||
return {"show_downloadable_models": self.show_downloadable_models}
|
||||
return {
|
||||
"show_downloadable_models": self.show_downloadable_models,
|
||||
"external_llm": self.external_llm.to_payload(),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, payload: object) -> SimpleSyrupSettings:
|
||||
@@ -50,7 +122,17 @@ class SimpleSyrupSettings:
|
||||
"SimpleSyrup settings payload is invalid. Expected "
|
||||
"show_downloadable_models to be a boolean."
|
||||
)
|
||||
return cls(show_downloadable_models=value)
|
||||
|
||||
try:
|
||||
external_llm = ExternalLLMSettings.from_payload(payload.get("external_llm"))
|
||||
except SimpleSyrupSettingsError as error:
|
||||
LOGGER.warning(
|
||||
"using default external llm settings after failed load",
|
||||
extra={"reason": str(error)},
|
||||
)
|
||||
external_llm = ExternalLLMSettings()
|
||||
|
||||
return cls(show_downloadable_models=value, external_llm=external_llm)
|
||||
|
||||
|
||||
class SimpleSyrupSettingsRepository:
|
||||
|
||||
@@ -60,7 +60,7 @@ class SettingsHandlers:
|
||||
|
||||
try:
|
||||
payload = await request.json()
|
||||
settings = SimpleSyrupSettings.from_payload(payload)
|
||||
settings = self._settings_from_request_payload(payload)
|
||||
except SimpleSyrupSettingsError as error:
|
||||
return web.json_response({"error": str(error)}, status=400)
|
||||
except Exception as error:
|
||||
@@ -76,6 +76,18 @@ class SettingsHandlers:
|
||||
saved = self._repository.save(settings)
|
||||
return web.json_response(saved.to_payload())
|
||||
|
||||
def _settings_from_request_payload(self, payload: object) -> SimpleSyrupSettings:
|
||||
"""Return validated settings while preserving omitted nested config."""
|
||||
|
||||
settings = SimpleSyrupSettings.from_payload(payload)
|
||||
if isinstance(payload, dict) and "external_llm" not in payload:
|
||||
current = self._repository.load()
|
||||
return SimpleSyrupSettings(
|
||||
show_downloadable_models=settings.show_downloadable_models,
|
||||
external_llm=current.external_llm,
|
||||
)
|
||||
return settings
|
||||
|
||||
|
||||
def register_settings_routes(
|
||||
repository: SimpleSyrupSettingsRepository | None = None,
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Build ComfyUI VAE option node expansions without owning VAE execution."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import Any, TypedDict, cast
|
||||
|
||||
|
||||
class ExpansionResult(TypedDict):
|
||||
"""ComfyUI dynamic expansion result with one output link."""
|
||||
|
||||
expand: dict[str, dict[str, Any]]
|
||||
result: tuple[list[object], ...]
|
||||
|
||||
|
||||
class VAEOptionsGraphBuilder:
|
||||
"""Expand VAE option choices to ComfyUI's native VAE nodes."""
|
||||
|
||||
def build_encode(
|
||||
self,
|
||||
*,
|
||||
pixels: object,
|
||||
vae: object,
|
||||
use_tiling: bool,
|
||||
tile_size: int,
|
||||
overlap: int,
|
||||
temporal_size: int,
|
||||
temporal_overlap: int,
|
||||
) -> ExpansionResult:
|
||||
"""Build a native VAE encode or tiled VAE encode expansion."""
|
||||
|
||||
inputs: dict[str, object] = {
|
||||
"pixels": _graph_value(pixels),
|
||||
"vae": _graph_value(vae),
|
||||
}
|
||||
class_type = "VAEEncode"
|
||||
if use_tiling:
|
||||
class_type = "VAEEncodeTiled"
|
||||
inputs.update(
|
||||
{
|
||||
"tile_size": int(tile_size),
|
||||
"overlap": int(overlap),
|
||||
"temporal_size": int(temporal_size),
|
||||
"temporal_overlap": int(temporal_overlap),
|
||||
}
|
||||
)
|
||||
|
||||
return _single_node_expansion(class_type, inputs)
|
||||
|
||||
def build_decode(
|
||||
self,
|
||||
*,
|
||||
samples: object,
|
||||
vae: object,
|
||||
use_tiling: bool,
|
||||
tile_size: int,
|
||||
overlap: int,
|
||||
temporal_size: int,
|
||||
temporal_overlap: int,
|
||||
) -> ExpansionResult:
|
||||
"""Build a native VAE decode or tiled VAE decode expansion."""
|
||||
|
||||
inputs: dict[str, object] = {
|
||||
"samples": _graph_value(samples),
|
||||
"vae": _graph_value(vae),
|
||||
}
|
||||
class_type = "VAEDecode"
|
||||
if use_tiling:
|
||||
class_type = "VAEDecodeTiled"
|
||||
inputs.update(
|
||||
{
|
||||
"tile_size": int(tile_size),
|
||||
"overlap": int(overlap),
|
||||
"temporal_size": int(temporal_size),
|
||||
"temporal_overlap": int(temporal_overlap),
|
||||
}
|
||||
)
|
||||
|
||||
return _single_node_expansion(class_type, inputs)
|
||||
|
||||
|
||||
def _single_node_expansion(
|
||||
class_type: str,
|
||||
inputs: dict[str, object],
|
||||
) -> ExpansionResult:
|
||||
"""Return a one-node dynamic expansion for a native ComfyUI node."""
|
||||
|
||||
builder = _graph_builder()
|
||||
node = builder.node(class_type, **inputs)
|
||||
return {
|
||||
"expand": cast(dict[str, dict[str, Any]], builder.finalize()),
|
||||
"result": (cast(list[object], node.out(0)),),
|
||||
}
|
||||
|
||||
|
||||
def _graph_value(value: object) -> object:
|
||||
"""Normalize tuple graph links to ComfyUI's serialized list shape."""
|
||||
|
||||
if isinstance(value, tuple) and len(value) == 2:
|
||||
node_id, output_slot = value
|
||||
if isinstance(node_id, str) and isinstance(output_slot, int):
|
||||
return [node_id, output_slot]
|
||||
return value
|
||||
|
||||
|
||||
def _graph_builder() -> Any:
|
||||
"""Return ComfyUI's graph builder without import-time Comfy coupling."""
|
||||
|
||||
graph_utils = import_module("comfy_execution.graph_utils")
|
||||
graph_builder = cast(Any, graph_utils.GraphBuilder)
|
||||
return graph_builder()
|
||||
@@ -64,12 +64,10 @@ class RegionalDetailSamplingBoundary(Protocol):
|
||||
denoise: float,
|
||||
global_prompt_weight: float,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample one full latent with paired regional conditioning."""
|
||||
|
||||
def apply_differential_diffusion(self, model: Any) -> Any:
|
||||
"""Return a model patched for feathered denoise masks."""
|
||||
|
||||
|
||||
class RegionalDetailResizeBoundary(Protocol):
|
||||
"""Image resize boundary used by regional detailing."""
|
||||
@@ -133,6 +131,7 @@ class RegionalDetailSampler:
|
||||
denoise: float,
|
||||
global_prompt_weight: float,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample one full latent with regional MultiDiffusion."""
|
||||
|
||||
@@ -150,13 +149,9 @@ class RegionalDetailSampler:
|
||||
denoise=denoise,
|
||||
global_prompt_weight=global_prompt_weight,
|
||||
preview_context=preview_context,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
|
||||
def apply_differential_diffusion(self, model: Any) -> Any:
|
||||
"""Patch a model for feathered denoise masks when ComfyUI supports it."""
|
||||
|
||||
return self._detail_sampler.apply_differential_diffusion(model)
|
||||
|
||||
|
||||
class DetailSEGSAsRegionsService:
|
||||
"""Detail provided SEGS through one regional MultiDiffusion pass."""
|
||||
@@ -273,12 +268,10 @@ class DetailSEGSAsRegionsService:
|
||||
latent_for_sampling = (
|
||||
self._with_noise_mask(latent, latent_regions) if noise_mask else latent
|
||||
)
|
||||
sampling_model = model
|
||||
if noise_mask and noise_mask_feather > 0:
|
||||
sampling_model = self._sampler.apply_differential_diffusion(model)
|
||||
differential_diffusion = noise_mask and noise_mask_feather > 0
|
||||
|
||||
sampled = self._sampler.sample_regions(
|
||||
model=sampling_model,
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
@@ -296,6 +289,7 @@ class DetailSEGSAsRegionsService:
|
||||
work_mask=image_union_mask,
|
||||
sampled_region=CropRegion(0, 0, image_width, image_height),
|
||||
),
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
decoded = self._sampler.decode(vae, sampled, tiled_decode)
|
||||
if decoded.shape[1:3] != image_tensor.shape[1:3]:
|
||||
|
||||
@@ -52,6 +52,7 @@ class TiledDiffusionLatentSamplingBoundary(Protocol):
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample a latent using the selected tiled diffusion mode."""
|
||||
|
||||
@@ -84,12 +85,10 @@ class TiledDetailSamplingBoundary(Protocol):
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample one latent crop with the requested tiled diffusion mode."""
|
||||
|
||||
def apply_differential_diffusion(self, model: Any) -> Any:
|
||||
"""Return a model patched for feathered denoise masks."""
|
||||
|
||||
|
||||
class TiledDetailResizeBoundary(Protocol):
|
||||
"""Image resize boundary used by tiled scale-factor detailing."""
|
||||
@@ -163,6 +162,7 @@ class TiledDetailSampler:
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample one latent crop with the selected tiled diffusion runtime."""
|
||||
|
||||
@@ -183,13 +183,9 @@ class TiledDetailSampler:
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=preview_context,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
|
||||
def apply_differential_diffusion(self, model: Any) -> Any:
|
||||
"""Patch a model for feathered denoise masks when ComfyUI supports it."""
|
||||
|
||||
return self._detail_sampler.apply_differential_diffusion(model)
|
||||
|
||||
|
||||
class DetailSEGSByScaleFactorTiledDiffusionService:
|
||||
"""Detail SEGS crops with tiled diffusion latent sampling."""
|
||||
@@ -251,9 +247,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
|
||||
return TiledDetailerResult(image=image_tensor.clone())
|
||||
|
||||
working_image = image_tensor.clone()
|
||||
sampling_model = model
|
||||
if noise_mask and noise_mask_feather > 0:
|
||||
sampling_model = self._sampler.apply_differential_diffusion(model)
|
||||
differential_diffusion = noise_mask and noise_mask_feather > 0
|
||||
|
||||
for index, segment in enumerate(segments):
|
||||
plan = build_detail_scale_plan(
|
||||
@@ -278,7 +272,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
|
||||
working_image=working_image,
|
||||
segment=segment,
|
||||
plan=plan,
|
||||
model=sampling_model,
|
||||
model=model,
|
||||
vae=vae,
|
||||
positive=select_conditioning(positive, index),
|
||||
negative=select_conditioning(negative, index),
|
||||
@@ -299,6 +293,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
|
||||
latent_tile_height=latent_tile_height,
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
|
||||
LOGGER.info(
|
||||
@@ -346,6 +341,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
|
||||
latent_tile_height: int,
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
differential_diffusion: bool,
|
||||
) -> torch.Tensor:
|
||||
"""Detail one segment with tiled diffusion and return the updated image."""
|
||||
|
||||
@@ -380,6 +376,7 @@ class DetailSEGSByScaleFactorTiledDiffusionService:
|
||||
latent_tile_height=latent_tile_height,
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
differential_diffusion=differential_diffusion,
|
||||
preview_context=DetailPreviewContext(
|
||||
image=working_image,
|
||||
work_region=segment.crop_region,
|
||||
|
||||
@@ -0,0 +1,316 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Application service for external LLM prompt generation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import replace
|
||||
from typing import Protocol
|
||||
|
||||
from ..domain.external_llm import (
|
||||
DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
ExternalLLMChatRequest,
|
||||
ExternalLLMChatResponse,
|
||||
ExternalLLMConfigError,
|
||||
ExternalLLMProviderError,
|
||||
normalize_base_url,
|
||||
)
|
||||
from ..runtime.external_llm_client import ExternalLLMClient
|
||||
from ..runtime.external_llm_images import ExternalLLMImageEncoder
|
||||
from ..runtime.external_llm_keyring import ExternalLLMKeyringStore
|
||||
from ..runtime.settings import (
|
||||
ExternalLLMSettings,
|
||||
SimpleSyrupSettings,
|
||||
SimpleSyrupSettingsRepository,
|
||||
)
|
||||
|
||||
CONFIGURE_EXTERNAL_LLM = "Configure external LLM endpoint"
|
||||
|
||||
|
||||
class SettingsRepository(Protocol):
|
||||
"""Settings persistence boundary used by external LLM services."""
|
||||
|
||||
def load(self) -> SimpleSyrupSettings:
|
||||
"""Load current SimpleSyrup settings."""
|
||||
|
||||
def save(self, settings: SimpleSyrupSettings) -> SimpleSyrupSettings:
|
||||
"""Persist current SimpleSyrup settings."""
|
||||
|
||||
|
||||
class KeyStore(Protocol):
|
||||
"""Credential storage boundary used by external LLM services."""
|
||||
|
||||
def has_api_key(self, base_url: str) -> bool:
|
||||
"""Return whether an API key exists for an endpoint."""
|
||||
|
||||
def get_api_key(self, base_url: str) -> str:
|
||||
"""Return an API key for an endpoint."""
|
||||
|
||||
def save_api_key(self, base_url: str, api_key: str) -> None:
|
||||
"""Store an API key for an endpoint."""
|
||||
|
||||
def delete_api_key(self, base_url: str) -> None:
|
||||
"""Delete an API key for an endpoint."""
|
||||
|
||||
|
||||
class ProviderClient(Protocol):
|
||||
"""External LLM provider client boundary."""
|
||||
|
||||
def list_models(self, base_url: str, api_key: str) -> tuple[str, ...]:
|
||||
"""Return provider model ids."""
|
||||
|
||||
def create_chat_completion(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
request: ExternalLLMChatRequest,
|
||||
) -> ExternalLLMChatResponse:
|
||||
"""Return provider chat completion content wrapper."""
|
||||
|
||||
|
||||
class ImageEncoder(Protocol):
|
||||
"""Image encoding boundary for external LLM vision inputs."""
|
||||
|
||||
def encode_first_image_as_data_url(self, image: object) -> str:
|
||||
"""Return the first ComfyUI IMAGE tensor item as a data URL."""
|
||||
|
||||
|
||||
class ExternalLLMPromptService:
|
||||
"""Coordinate external LLM settings, credentials, and provider calls."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
settings_repository: SettingsRepository | None = None,
|
||||
key_store: KeyStore | None = None,
|
||||
client: ProviderClient | None = None,
|
||||
image_encoder: ImageEncoder | None = None,
|
||||
) -> None:
|
||||
"""Create the service with injectable runtime boundaries."""
|
||||
|
||||
self._settings_repository = (
|
||||
settings_repository or SimpleSyrupSettingsRepository()
|
||||
)
|
||||
self._key_store = key_store or ExternalLLMKeyringStore()
|
||||
self._client = client or ExternalLLMClient()
|
||||
self._image_encoder = image_encoder or ExternalLLMImageEncoder()
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return cached provider model choices for Comfy dropdowns."""
|
||||
|
||||
models = self._settings_repository.load().external_llm.cached_models
|
||||
return list(models) if models else [CONFIGURE_EXTERNAL_LLM]
|
||||
|
||||
def provider_is_configured(self) -> bool:
|
||||
"""Return whether endpoint and API key are configured."""
|
||||
|
||||
external = self._settings_repository.load().external_llm
|
||||
return bool(external.base_url) and self._key_store.has_api_key(
|
||||
external.base_url
|
||||
)
|
||||
|
||||
def settings_payload(self) -> dict[str, object]:
|
||||
"""Return non-secret external LLM settings for frontend routes."""
|
||||
|
||||
external = self._settings_repository.load().external_llm
|
||||
has_api_key = (
|
||||
self._key_store.has_api_key(external.base_url)
|
||||
if external.base_url
|
||||
else False
|
||||
)
|
||||
return {
|
||||
"base_url": external.base_url,
|
||||
"cached_models": list(external.cached_models),
|
||||
"default_model": external.default_model,
|
||||
"has_api_key": has_api_key,
|
||||
}
|
||||
|
||||
def refresh_models(self) -> tuple[str, ...]:
|
||||
"""Refresh cached provider models when endpoint credentials exist."""
|
||||
|
||||
settings = self._settings_repository.load()
|
||||
external = settings.external_llm
|
||||
if not external.base_url:
|
||||
return external.cached_models
|
||||
|
||||
if not self._key_store.has_api_key(external.base_url):
|
||||
return external.cached_models
|
||||
|
||||
api_key = self._key_store.get_api_key(external.base_url)
|
||||
models = self._client.list_models(external.base_url, api_key)
|
||||
default_model = external.default_model
|
||||
if models and default_model not in models:
|
||||
default_model = models[0]
|
||||
if not models:
|
||||
default_model = ""
|
||||
|
||||
saved = self._settings_repository.save(
|
||||
replace(
|
||||
settings,
|
||||
external_llm=replace(
|
||||
external,
|
||||
cached_models=models,
|
||||
default_model=default_model,
|
||||
),
|
||||
)
|
||||
)
|
||||
return saved.external_llm.cached_models
|
||||
|
||||
def save_config(
|
||||
self, base_url: str, default_model: str = ""
|
||||
) -> SimpleSyrupSettings:
|
||||
"""Persist non-secret external LLM endpoint settings."""
|
||||
|
||||
settings = self._settings_repository.load()
|
||||
normalized_url = normalize_base_url(base_url) if base_url.strip() else ""
|
||||
endpoint_changed = normalized_url != settings.external_llm.base_url
|
||||
cached_models = () if endpoint_changed else settings.external_llm.cached_models
|
||||
default = default_model.strip()
|
||||
if endpoint_changed:
|
||||
default = ""
|
||||
if default and cached_models and default not in cached_models:
|
||||
raise ExternalLLMConfigError(
|
||||
"Default external LLM model must be one of the cached models."
|
||||
)
|
||||
|
||||
self._settings_repository.save(
|
||||
replace(
|
||||
settings,
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url=normalized_url,
|
||||
cached_models=cached_models,
|
||||
default_model=default,
|
||||
),
|
||||
)
|
||||
)
|
||||
try:
|
||||
self.refresh_models()
|
||||
except ExternalLLMProviderError:
|
||||
pass
|
||||
return self._settings_repository.load()
|
||||
|
||||
def save_api_key(self, api_key: str) -> SimpleSyrupSettings:
|
||||
"""Store the API key for the configured endpoint."""
|
||||
|
||||
settings = self._settings_repository.load()
|
||||
base_url = settings.external_llm.base_url
|
||||
if not base_url:
|
||||
raise ExternalLLMConfigError(
|
||||
"Configure an external LLM endpoint in SimpleSyrup settings before "
|
||||
"adding an API key."
|
||||
)
|
||||
self._key_store.save_api_key(base_url, api_key)
|
||||
try:
|
||||
self.refresh_models()
|
||||
except ExternalLLMProviderError:
|
||||
pass
|
||||
return self._settings_repository.load()
|
||||
|
||||
def delete_api_key(self) -> SimpleSyrupSettings:
|
||||
"""Delete the API key for the configured endpoint."""
|
||||
|
||||
settings = self._settings_repository.load()
|
||||
base_url = settings.external_llm.base_url
|
||||
if base_url:
|
||||
self._key_store.delete_api_key(base_url)
|
||||
return settings
|
||||
|
||||
def generate(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
image: object | None = None,
|
||||
) -> str:
|
||||
"""Generate an assistant response for the supplied prompt pair."""
|
||||
|
||||
return self.generate_with_image_data_url(
|
||||
model=model,
|
||||
system_prompt=system_prompt,
|
||||
user_prompt=user_prompt,
|
||||
max_tokens=max_tokens,
|
||||
reasoning_effort=reasoning_effort,
|
||||
image_data_url=(
|
||||
None
|
||||
if image is None
|
||||
else self._image_encoder.encode_first_image_as_data_url(image)
|
||||
),
|
||||
)
|
||||
|
||||
def generate_with_image_data_url(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
image_data_url: str | None = None,
|
||||
) -> str:
|
||||
"""Generate an assistant response with a pre-encoded optional image."""
|
||||
|
||||
selected_model = self._resolve_model_for_execution(model)
|
||||
|
||||
external = self._settings_repository.load().external_llm
|
||||
if not external.base_url:
|
||||
raise ExternalLLMConfigError(
|
||||
"Configure an external LLM endpoint in SimpleSyrup settings before "
|
||||
"using this node."
|
||||
)
|
||||
|
||||
api_key = self._key_store.get_api_key(external.base_url)
|
||||
if not api_key:
|
||||
raise ExternalLLMConfigError(
|
||||
"Add an API key for the configured external LLM endpoint in "
|
||||
"SimpleSyrup settings."
|
||||
)
|
||||
|
||||
request = ExternalLLMChatRequest(
|
||||
model=selected_model,
|
||||
system_prompt=system_prompt,
|
||||
user_prompt=user_prompt,
|
||||
max_tokens=max_tokens,
|
||||
reasoning_effort=reasoning_effort,
|
||||
image_data_url=image_data_url,
|
||||
)
|
||||
response = self._client.create_chat_completion(
|
||||
external.base_url,
|
||||
api_key,
|
||||
request,
|
||||
)
|
||||
return response.content
|
||||
|
||||
def _resolve_model_for_execution(self, model: str) -> str:
|
||||
"""Return a concrete provider model, refreshing stale sentinel choices."""
|
||||
|
||||
selected_model = model.strip()
|
||||
if selected_model != CONFIGURE_EXTERNAL_LLM:
|
||||
return selected_model
|
||||
|
||||
settings = self._settings_repository.load()
|
||||
external = settings.external_llm
|
||||
if not external.base_url:
|
||||
raise ExternalLLMConfigError(
|
||||
"Configure an external LLM endpoint in SimpleSyrup settings before "
|
||||
"using this node."
|
||||
)
|
||||
if not self._key_store.has_api_key(external.base_url):
|
||||
raise ExternalLLMConfigError(
|
||||
"Add an API key for the configured external LLM endpoint in "
|
||||
"SimpleSyrup settings."
|
||||
)
|
||||
|
||||
models = external.cached_models or self.refresh_models()
|
||||
refreshed_external = self._settings_repository.load().external_llm
|
||||
default_model = refreshed_external.default_model
|
||||
if default_model and default_model in models:
|
||||
return default_model
|
||||
if models:
|
||||
return models[0]
|
||||
raise ExternalLLMConfigError(
|
||||
"The configured external LLM endpoint did not report any models. "
|
||||
"Check the endpoint URL and API key in SimpleSyrup settings."
|
||||
)
|
||||
@@ -0,0 +1,295 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Application service for external-LLM tagging of existing SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.conditioning_batch import ConditioningBatch
|
||||
from ..domain.external_llm import (
|
||||
DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
)
|
||||
from ..domain.prompt_composition import prefix_prompt
|
||||
from ..domain.segs import (
|
||||
ImpactSegs,
|
||||
NativeSegs,
|
||||
Segment,
|
||||
coerce_segs,
|
||||
to_impact_compatible_segs,
|
||||
)
|
||||
from ..masking.segs_mask_ops import validate_single_image
|
||||
from ..runtime.conditioning_encoding import ComfyConditioningEncoder
|
||||
from ..runtime.external_llm_images import ExternalLLMSegsImageEncoder
|
||||
from ..runtime.progress import ProgressReporter, create_comfy_progress
|
||||
from ..shared.logging import get_logger
|
||||
from .external_llm_prompt_service import ExternalLLMPromptService
|
||||
|
||||
LOGGER = get_logger(__name__)
|
||||
OPERATION = "Tag SEGS w/ External LLM"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LLMTagFormattingControls:
|
||||
"""Store formatting controls for external LLM tag responses."""
|
||||
|
||||
replace_underscore: bool = True
|
||||
trailing_comma: bool = False
|
||||
exclude_tags: str = ""
|
||||
|
||||
|
||||
class ExternalLLMGenerationBoundary(Protocol):
|
||||
"""External LLM provider execution boundary."""
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return cached provider model choices for Comfy dropdowns."""
|
||||
|
||||
def generate_with_image_data_url(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
image_data_url: str | None = None,
|
||||
) -> str:
|
||||
"""Return one assistant response for a pre-encoded image."""
|
||||
|
||||
|
||||
class SegmentImageEncodingBoundary(Protocol):
|
||||
"""Encode ordered SEG crops as LLM vision images."""
|
||||
|
||||
def encode_segment_as_data_url(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
segment: Segment,
|
||||
mode: str,
|
||||
) -> str:
|
||||
"""Return one SEG crop image data URL."""
|
||||
|
||||
|
||||
class ConditioningEncodingBoundary(Protocol):
|
||||
"""Encode ordered prompts into a conditioning batch."""
|
||||
|
||||
def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch:
|
||||
"""Return conditioning entries in prompt order."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TagSEGSWithExternalLLMResult:
|
||||
"""Return unchanged SEGS and aligned external-LLM conditioning."""
|
||||
|
||||
segs: ImpactSegs
|
||||
positive: ConditioningBatch
|
||||
|
||||
|
||||
class TagSEGSWithExternalLLMService:
|
||||
"""Caption provided SEGS crops with an external LLM and encode conditioning."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
llm: ExternalLLMGenerationBoundary | None = None,
|
||||
image_encoder: SegmentImageEncodingBoundary | None = None,
|
||||
conditioning_encoder: ConditioningEncodingBoundary | None = None,
|
||||
progress_factory: Callable[[int], ProgressReporter] | None = None,
|
||||
) -> None:
|
||||
"""Create the service with injectable runtime boundaries."""
|
||||
|
||||
self._llm = llm or ExternalLLMPromptService()
|
||||
self._image_encoder = image_encoder or ExternalLLMSegsImageEncoder()
|
||||
self._conditioning_encoder = conditioning_encoder or ComfyConditioningEncoder()
|
||||
self._progress_factory = progress_factory or create_comfy_progress
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return cached provider model choices for Comfy dropdowns."""
|
||||
|
||||
return self._llm.model_choices()
|
||||
|
||||
def tag(
|
||||
self,
|
||||
image: object,
|
||||
segs: object,
|
||||
clip: Any,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
universal_positive: str,
|
||||
seg_image_mode: str,
|
||||
formatting: LLMTagFormattingControls,
|
||||
max_tokens: int = DEFAULT_EXTERNAL_LLM_MAX_TOKENS,
|
||||
reasoning_effort: str = DEFAULT_EXTERNAL_LLM_REASONING_EFFORT,
|
||||
) -> TagSEGSWithExternalLLMResult:
|
||||
"""Return original SEGS plus LLM-derived conditioning in segment order."""
|
||||
|
||||
image_tensor = validate_single_image(image, OPERATION)
|
||||
native_segs = coerce_segs(segs)
|
||||
self._validate_segs_target_image(native_segs, image_tensor)
|
||||
_header, segments = native_segs
|
||||
if not segments:
|
||||
raise ValueError("No SEGS were provided for Tag SEGS w/ External LLM.")
|
||||
|
||||
progress = self._progress_factory(len(segments) + 2)
|
||||
progress.update(1)
|
||||
prompts: list[str] = []
|
||||
for segment in segments:
|
||||
prompts.append(
|
||||
self._prompt_for_segment(
|
||||
image=image_tensor,
|
||||
segment=segment,
|
||||
model=model,
|
||||
system_prompt=system_prompt,
|
||||
user_prompt=user_prompt,
|
||||
universal_positive=universal_positive,
|
||||
seg_image_mode=seg_image_mode,
|
||||
formatting=formatting,
|
||||
max_tokens=max_tokens,
|
||||
reasoning_effort=reasoning_effort,
|
||||
)
|
||||
)
|
||||
progress.update(1)
|
||||
|
||||
positive = self._conditioning_encoder.encode_batch(clip, tuple(prompts))
|
||||
progress.update(1)
|
||||
if len(positive.entries) != len(segments):
|
||||
raise ValueError(
|
||||
"Conditioning encoder returned "
|
||||
f"{len(positive.entries)} entries for {len(segments)} SEGS."
|
||||
)
|
||||
|
||||
LOGGER.info(
|
||||
"Tag SEGS w/ External LLM pass completed",
|
||||
extra={
|
||||
"operation": "tag_segs_with_external_llm",
|
||||
"segment_count": len(segments),
|
||||
"external_llm_model": model,
|
||||
"seg_image_mode": seg_image_mode,
|
||||
"universal_positive_present": bool(universal_positive.strip()),
|
||||
"replace_underscore": formatting.replace_underscore,
|
||||
"trailing_comma": formatting.trailing_comma,
|
||||
"exclude_tags_present": bool(formatting.exclude_tags.strip()),
|
||||
},
|
||||
)
|
||||
return TagSEGSWithExternalLLMResult(
|
||||
segs=to_impact_compatible_segs(native_segs),
|
||||
positive=positive,
|
||||
)
|
||||
|
||||
def _prompt_for_segment(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
segment: Segment,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
universal_positive: str,
|
||||
seg_image_mode: str,
|
||||
formatting: LLMTagFormattingControls,
|
||||
max_tokens: int,
|
||||
reasoning_effort: str,
|
||||
) -> str:
|
||||
"""Generate, format, and prefix one segment prompt."""
|
||||
|
||||
image_data_url = self._image_encoder.encode_segment_as_data_url(
|
||||
image,
|
||||
segment,
|
||||
seg_image_mode,
|
||||
)
|
||||
response = self._llm.generate_with_image_data_url(
|
||||
model=model,
|
||||
system_prompt=system_prompt,
|
||||
user_prompt=user_prompt,
|
||||
max_tokens=max_tokens,
|
||||
reasoning_effort=reasoning_effort,
|
||||
image_data_url=image_data_url,
|
||||
)
|
||||
prompt = format_external_llm_tags(response, formatting)
|
||||
return prefix_prompt(universal_positive, prompt)
|
||||
|
||||
def _validate_segs_target_image(
|
||||
self,
|
||||
segs: NativeSegs,
|
||||
image: torch.Tensor,
|
||||
) -> None:
|
||||
"""Reject SEGS that cannot be cropped from the provided image."""
|
||||
|
||||
header, segments = segs
|
||||
image_height = int(image.shape[1])
|
||||
image_width = int(image.shape[2])
|
||||
if header != (image_height, image_width):
|
||||
raise ValueError(
|
||||
f"{OPERATION} requires SEGS header dimensions to match the image: "
|
||||
f"SEGS is {header[0]}x{header[1]}, image is "
|
||||
f"{image_height}x{image_width}."
|
||||
)
|
||||
for index, segment in enumerate(segments):
|
||||
_validate_segment_crop(segment, index, image_height, image_width)
|
||||
|
||||
|
||||
def format_external_llm_tags(
|
||||
response: str,
|
||||
controls: LLMTagFormattingControls,
|
||||
) -> str:
|
||||
"""Format one external LLM response into comma-separated prompt tags."""
|
||||
|
||||
stripped = response.strip()
|
||||
if not stripped:
|
||||
raise ValueError("External LLM returned an empty response for a SEG.")
|
||||
|
||||
excluded = _excluded_tags(controls)
|
||||
tags: list[str] = []
|
||||
for chunk in stripped.split(","):
|
||||
tag = _normalize_tag(chunk, controls.replace_underscore)
|
||||
if not tag or tag.lower() in excluded:
|
||||
continue
|
||||
tags.append(tag)
|
||||
if not tags:
|
||||
raise ValueError("External LLM response had no usable tags after exclusions.")
|
||||
|
||||
prompt = ", ".join(tags)
|
||||
if controls.trailing_comma and not prompt.endswith(","):
|
||||
return f"{prompt},"
|
||||
return prompt
|
||||
|
||||
|
||||
def _excluded_tags(controls: LLMTagFormattingControls) -> set[str]:
|
||||
"""Return normalized excluded tag names."""
|
||||
|
||||
return {
|
||||
tag.lower()
|
||||
for raw_tag in controls.exclude_tags.split(",")
|
||||
if (tag := _normalize_tag(raw_tag, controls.replace_underscore))
|
||||
}
|
||||
|
||||
|
||||
def _normalize_tag(value: str, replace_underscore: bool) -> str:
|
||||
"""Return the prompt-facing form of one tag-like response chunk."""
|
||||
|
||||
tag = value.strip()
|
||||
if replace_underscore:
|
||||
return tag.replace("_", " ")
|
||||
return tag
|
||||
|
||||
|
||||
def _validate_segment_crop(
|
||||
segment: Segment,
|
||||
index: int,
|
||||
image_height: int,
|
||||
image_width: int,
|
||||
) -> None:
|
||||
"""Reject a segment crop region that falls outside the image."""
|
||||
|
||||
region = segment.crop_region
|
||||
if region.right <= image_width and region.bottom <= image_height:
|
||||
return
|
||||
raise ValueError(
|
||||
f"{OPERATION} SEG {index} ('{segment.label}') crop_region must fit "
|
||||
f"inside the image; got ({region.left}, {region.top}, {region.right}, "
|
||||
f"{region.bottom}) for image {image_height}x{image_width}."
|
||||
)
|
||||
@@ -0,0 +1,173 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Application service for WD14 tagging of existing SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.conditioning_batch import ConditioningBatch
|
||||
from ..domain.prompt_composition import prefix_prompt
|
||||
from ..domain.segs import (
|
||||
ImpactSegs,
|
||||
NativeSegs,
|
||||
Segment,
|
||||
coerce_segs,
|
||||
to_impact_compatible_segs,
|
||||
)
|
||||
from ..masking.segs_mask_ops import crop_image, validate_single_image
|
||||
from ..runtime.conditioning_encoding import ComfyConditioningEncoder
|
||||
from ..runtime.loaded_models import LoadedWD14Tagger, unwrap_wd14_tagger
|
||||
from ..runtime.progress import ProgressReporter, create_comfy_progress
|
||||
from ..runtime.wd14_tagger import WD14TagFormattingControls, WD14Tagger
|
||||
from ..shared.logging import get_logger
|
||||
|
||||
LOGGER = get_logger(__name__)
|
||||
OPERATION = "Tag SEGS w/ WD14"
|
||||
|
||||
|
||||
class WD14TaggingBoundary(Protocol):
|
||||
"""Tag ordered image crops."""
|
||||
|
||||
def tag_images(
|
||||
self,
|
||||
loaded_tagger: LoadedWD14Tagger,
|
||||
images: tuple[torch.Tensor, ...],
|
||||
controls: WD14TagFormattingControls,
|
||||
progress: ProgressReporter | None = None,
|
||||
) -> tuple[str, ...]:
|
||||
"""Return one tag string per image in input order."""
|
||||
|
||||
|
||||
class ConditioningEncodingBoundary(Protocol):
|
||||
"""Encode ordered prompts into a conditioning batch."""
|
||||
|
||||
def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch:
|
||||
"""Return conditioning entries in prompt order."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TagSEGSWithWD14Result:
|
||||
"""Return unchanged SEGS and aligned positive conditioning."""
|
||||
|
||||
segs: ImpactSegs
|
||||
positive: ConditioningBatch
|
||||
|
||||
|
||||
class TagSEGSWithWD14Service:
|
||||
"""Tag provided SEGS crops and encode aligned regional conditioning."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tagger: WD14TaggingBoundary | None = None,
|
||||
encoder: ConditioningEncodingBoundary | None = None,
|
||||
progress_factory: Callable[[int], ProgressReporter] | None = None,
|
||||
) -> None:
|
||||
"""Create the service with injectable collaborators for tests."""
|
||||
|
||||
self._tagger = tagger or WD14Tagger()
|
||||
self._encoder = encoder or ComfyConditioningEncoder()
|
||||
self._progress_factory = progress_factory or create_comfy_progress
|
||||
|
||||
def tag(
|
||||
self,
|
||||
image: object,
|
||||
segs: object,
|
||||
clip: Any,
|
||||
wd14_tagger: object,
|
||||
tag_controls: WD14TagFormattingControls,
|
||||
universal_positive: str,
|
||||
) -> TagSEGSWithWD14Result:
|
||||
"""Return original SEGS plus WD14-derived conditioning in segment order."""
|
||||
|
||||
image_tensor = validate_single_image(image, OPERATION)
|
||||
native_segs = coerce_segs(segs)
|
||||
self._validate_segs_target_image(native_segs, image_tensor)
|
||||
_header, segments = native_segs
|
||||
if not segments:
|
||||
raise ValueError("No SEGS were provided for Tag SEGS w/ WD14.")
|
||||
|
||||
loaded_tagger = unwrap_wd14_tagger(wd14_tagger)
|
||||
progress = self._progress_factory(len(segments) + 2)
|
||||
progress.update(1)
|
||||
crops = tuple(
|
||||
crop_image(image_tensor, segment.crop_region) for segment in segments
|
||||
)
|
||||
tags = self._tagger.tag_images(
|
||||
loaded_tagger,
|
||||
crops,
|
||||
tag_controls,
|
||||
progress=progress,
|
||||
)
|
||||
if len(tags) != len(segments):
|
||||
raise ValueError(
|
||||
f"WD14 tagger returned {len(tags)} tag(s) for {len(segments)} SEGS."
|
||||
)
|
||||
|
||||
prompts = tuple(prefix_prompt(universal_positive, tag) for tag in tags)
|
||||
positive = self._encoder.encode_batch(clip, prompts)
|
||||
progress.update(1)
|
||||
if len(positive.entries) != len(segments):
|
||||
raise ValueError(
|
||||
"Conditioning encoder returned "
|
||||
f"{len(positive.entries)} entries for {len(segments)} SEGS."
|
||||
)
|
||||
|
||||
LOGGER.info(
|
||||
"Tag SEGS w/ WD14 pass completed",
|
||||
extra={
|
||||
"operation": "tag_segs_with_wd14",
|
||||
"segment_count": len(segments),
|
||||
"wd14_model": loaded_tagger.model_id,
|
||||
"threshold": tag_controls.threshold,
|
||||
"character_threshold": tag_controls.character_threshold,
|
||||
"universal_positive_present": bool(universal_positive.strip()),
|
||||
},
|
||||
)
|
||||
return TagSEGSWithWD14Result(
|
||||
segs=to_impact_compatible_segs(native_segs),
|
||||
positive=positive,
|
||||
)
|
||||
|
||||
def _validate_segs_target_image(
|
||||
self,
|
||||
segs: NativeSegs,
|
||||
image: torch.Tensor,
|
||||
) -> None:
|
||||
"""Reject SEGS that cannot be cropped from the provided image."""
|
||||
|
||||
header, segments = segs
|
||||
image_height = int(image.shape[1])
|
||||
image_width = int(image.shape[2])
|
||||
if header != (image_height, image_width):
|
||||
raise ValueError(
|
||||
f"{OPERATION} requires SEGS header dimensions to match the image: "
|
||||
f"SEGS is {header[0]}x{header[1]}, image is "
|
||||
f"{image_height}x{image_width}."
|
||||
)
|
||||
for index, segment in enumerate(segments):
|
||||
_validate_segment_crop(segment, index, image_height, image_width)
|
||||
|
||||
|
||||
def _validate_segment_crop(
|
||||
segment: Segment,
|
||||
index: int,
|
||||
image_height: int,
|
||||
image_width: int,
|
||||
) -> None:
|
||||
"""Reject a segment crop region that falls outside the image."""
|
||||
|
||||
region = segment.crop_region
|
||||
if region.right <= image_width and region.bottom <= image_height:
|
||||
return
|
||||
raise ValueError(
|
||||
f"{OPERATION} SEG {index} ('{segment.label}') crop_region must fit "
|
||||
f"inside the image; got ({region.left}, {region.top}, {region.right}, "
|
||||
f"{region.bottom}) for image {image_height}x{image_width}."
|
||||
)
|
||||
@@ -8,6 +8,9 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from ..domain.conditioning_batch import ConditioningBatch, select_conditioning
|
||||
from ..domain.tiled_diffusion import validate_tiled_diffusion_mode
|
||||
from ..runtime import mixture_of_diffusers_sampling, multidiffusion_sampling
|
||||
from ..runtime.detail_previews import DetailPreviewContext
|
||||
@@ -37,10 +40,31 @@ class TiledDiffusionSamplingService:
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Sample a latent with the selected tiled diffusion method."""
|
||||
|
||||
validate_tiled_diffusion_mode(diffusion_mode)
|
||||
if self._uses_conditioning_batch(positive, negative):
|
||||
return self._sample_conditioning_batch(
|
||||
diffusion_mode=diffusion_mode,
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=positive,
|
||||
negative=negative,
|
||||
latent_image=latent_image,
|
||||
denoise=denoise,
|
||||
latent_tile_width=latent_tile_width,
|
||||
latent_tile_height=latent_tile_height,
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=preview_context,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
if diffusion_mode == "multidiffusion":
|
||||
return multidiffusion_sampling.sample_multidiffusion(
|
||||
model=model,
|
||||
@@ -58,6 +82,7 @@ class TiledDiffusionSamplingService:
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=preview_context,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
return mixture_of_diffusers_sampling.sample_mixture_of_diffusers(
|
||||
model=model,
|
||||
@@ -75,4 +100,106 @@ class TiledDiffusionSamplingService:
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=preview_context,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
|
||||
def _sample_conditioning_batch(
|
||||
self,
|
||||
*,
|
||||
diffusion_mode: str,
|
||||
model: Any,
|
||||
seed: int,
|
||||
steps: int,
|
||||
cfg: float,
|
||||
sampler_name: str,
|
||||
scheduler: str,
|
||||
positive: Any,
|
||||
negative: Any,
|
||||
latent_image: Latent,
|
||||
denoise: float,
|
||||
latent_tile_width: int,
|
||||
latent_tile_height: int,
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None,
|
||||
differential_diffusion: bool,
|
||||
) -> Latent:
|
||||
"""Sample latent batch items one at a time with selected conditioning."""
|
||||
|
||||
latent_samples = latent_image["samples"]
|
||||
if not isinstance(latent_samples, torch.Tensor):
|
||||
raise TypeError("Tiled diffusion latent samples must be a torch.Tensor.")
|
||||
|
||||
outputs: list[torch.Tensor] = []
|
||||
for index in range(int(latent_samples.shape[0])):
|
||||
item_latent = self._single_item_latent(latent_image, index)
|
||||
output = self.sample(
|
||||
diffusion_mode=diffusion_mode,
|
||||
model=model,
|
||||
seed=seed,
|
||||
steps=steps,
|
||||
cfg=cfg,
|
||||
sampler_name=sampler_name,
|
||||
scheduler=scheduler,
|
||||
positive=select_conditioning(positive, index),
|
||||
negative=select_conditioning(negative, index),
|
||||
latent_image=item_latent,
|
||||
denoise=denoise,
|
||||
latent_tile_width=latent_tile_width,
|
||||
latent_tile_height=latent_tile_height,
|
||||
latent_tile_overlap=latent_tile_overlap,
|
||||
latent_tile_batch_size=latent_tile_batch_size,
|
||||
preview_context=preview_context,
|
||||
differential_diffusion=differential_diffusion,
|
||||
)
|
||||
output_samples = output["samples"]
|
||||
if not isinstance(output_samples, torch.Tensor):
|
||||
raise TypeError(
|
||||
"Tiled diffusion output samples must be a torch.Tensor."
|
||||
)
|
||||
outputs.append(output_samples)
|
||||
|
||||
result = latent_image.copy()
|
||||
result.pop("downscale_ratio_spacial", None)
|
||||
result["samples"] = torch.cat(outputs, dim=0)
|
||||
return result
|
||||
|
||||
def _single_item_latent(self, latent_image: Latent, index: int) -> Latent:
|
||||
"""Return a latent dictionary for one batch item."""
|
||||
|
||||
latent_samples = latent_image["samples"]
|
||||
if not isinstance(latent_samples, torch.Tensor):
|
||||
raise TypeError("Tiled diffusion latent samples must be a torch.Tensor.")
|
||||
item = latent_image.copy()
|
||||
item["samples"] = latent_samples[index : index + 1]
|
||||
if "batch_index" in item:
|
||||
item["batch_index"] = [item["batch_index"][index]]
|
||||
if "noise_mask" in item:
|
||||
item["noise_mask"] = self._slice_noise_mask(
|
||||
item["noise_mask"],
|
||||
index,
|
||||
latent_samples,
|
||||
)
|
||||
return item
|
||||
|
||||
def _slice_noise_mask(
|
||||
self,
|
||||
noise_mask: Any,
|
||||
index: int,
|
||||
latent_samples: torch.Tensor,
|
||||
) -> Any:
|
||||
"""Return the noise mask slice matching one latent batch item."""
|
||||
|
||||
if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int(
|
||||
latent_samples.shape[0],
|
||||
):
|
||||
return noise_mask[index : index + 1]
|
||||
return noise_mask
|
||||
|
||||
def _uses_conditioning_batch(self, positive: Any, negative: Any) -> bool:
|
||||
"""Return whether tiled sampling needs per-item conditioning selection."""
|
||||
|
||||
return isinstance(positive, ConditioningBatch) or isinstance(
|
||||
negative,
|
||||
ConditioningBatch,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Batch Region Conditioning legacy node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.nodes.batch_region_conditioning import BatchRegionConditioning
|
||||
|
||||
|
||||
def test_batch_region_conditioning_contract() -> None:
|
||||
"""Batch Region Conditioning exposes a mixed-input legacy contract."""
|
||||
|
||||
inputs = BatchRegionConditioning.INPUT_TYPES()
|
||||
|
||||
assert BatchRegionConditioning.RETURN_TYPES == ("CONDITIONING_BATCH",)
|
||||
assert BatchRegionConditioning.RETURN_NAMES == ("batch",)
|
||||
assert BatchRegionConditioning.FUNCTION == "batch"
|
||||
assert BatchRegionConditioning.CATEGORY == "SimpleSyrup/Conditioning"
|
||||
assert list(inputs["required"]) == ["first", "second"]
|
||||
assert inputs["required"]["first"][0] == "CONDITIONING,CONDITIONING_BATCH"
|
||||
assert inputs["required"]["second"][0] == "CONDITIONING,CONDITIONING_BATCH"
|
||||
|
||||
|
||||
def test_batch_region_conditioning_node_flattens_mixed_inputs() -> None:
|
||||
"""The legacy node flattens batches and normal conditionings in order."""
|
||||
|
||||
auto = ConditioningBatch(("auto 1", "auto 2"))
|
||||
hand = "hand 1"
|
||||
|
||||
(batch,) = BatchRegionConditioning().batch(auto, hand)
|
||||
|
||||
assert batch == ConditioningBatch(("auto 1", "auto 2", "hand 1"))
|
||||
@@ -0,0 +1,50 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Batch Region Conditioning Comfy v3 wrapper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.nodes_v3.batch_region_conditioning import (
|
||||
BatchRegionConditioningV3,
|
||||
)
|
||||
|
||||
|
||||
def test_batch_region_conditioning_v3_schema_uses_mixed_autogrow_inputs() -> None:
|
||||
"""The v3 schema accepts conditioning values and conditioning batches."""
|
||||
|
||||
schema = BatchRegionConditioningV3.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.BatchRegionConditioning"
|
||||
assert schema.display_name == "Batch Region Conditioning"
|
||||
assert schema.category == "SimpleSyrup/Conditioning"
|
||||
assert [input_item.id for input_item in schema.inputs] == ["conditioning_inputs"]
|
||||
assert schema.inputs[0].io_type == "COMFY_AUTOGROW_V3"
|
||||
assert schema.inputs[0].template.prefix == "conditioning"
|
||||
assert schema.inputs[0].template.min == 2
|
||||
assert schema.inputs[0].template.max == 50
|
||||
assert schema.inputs[0].template.input.get_io_type() == (
|
||||
"CONDITIONING,CONDITIONING_BATCH"
|
||||
)
|
||||
assert [output.id for output in schema.outputs] == ["batch"]
|
||||
assert schema.outputs[0].io_type == "CONDITIONING_BATCH"
|
||||
|
||||
|
||||
def test_batch_region_conditioning_v3_execute_flattens_inputs() -> None:
|
||||
"""The v3 wrapper batches mixed inputs in Autogrow insertion order."""
|
||||
|
||||
auto = ConditioningBatch(("auto 1", "auto 2"))
|
||||
hand = "hand 1"
|
||||
extra = ConditioningBatch(("auto 3",))
|
||||
|
||||
(batch,) = BatchRegionConditioningV3.execute(
|
||||
{
|
||||
"conditioning0": auto,
|
||||
"conditioning1": hand,
|
||||
"conditioning2": extra,
|
||||
}
|
||||
)
|
||||
|
||||
assert batch == ConditioningBatch(("auto 1", "auto 2", "hand 1", "auto 3"))
|
||||
@@ -0,0 +1,53 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Batch SEGS legacy node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import cast
|
||||
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, ImpactSegs, Segment
|
||||
from simple_syrup.nodes.batch_segs import BatchSEGS
|
||||
|
||||
|
||||
def test_batch_segs_contract() -> None:
|
||||
"""Batch SEGS exposes a legacy two-input chainable contract."""
|
||||
|
||||
inputs = BatchSEGS.INPUT_TYPES()
|
||||
|
||||
assert BatchSEGS.RETURN_TYPES == ("SEGS",)
|
||||
assert BatchSEGS.RETURN_NAMES == ("segs",)
|
||||
assert BatchSEGS.FUNCTION == "batch"
|
||||
assert BatchSEGS.CATEGORY == "SimpleSyrup/Detection"
|
||||
assert list(inputs["required"]) == ["first", "second"]
|
||||
assert inputs["required"]["first"][0] == "SEGS"
|
||||
assert inputs["required"]["second"][0] == "SEGS"
|
||||
|
||||
|
||||
def test_batch_segs_node_batches_in_input_order() -> None:
|
||||
"""The legacy node returns Impact-compatible batched SEGS."""
|
||||
|
||||
first = ((8, 8), [_segment("1"), _segment("2")])
|
||||
second = ((8, 8), [_segment("3")])
|
||||
|
||||
(raw_segs,) = BatchSEGS().batch(first, second)
|
||||
segs = cast(ImpactSegs, raw_segs)
|
||||
|
||||
_header, segments = segs
|
||||
assert isinstance(segments, list)
|
||||
assert [segment.label for segment in segments] == ["1", "2", "3"]
|
||||
|
||||
|
||||
def _segment(label: str) -> Segment:
|
||||
"""Create a small test segment."""
|
||||
|
||||
return Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask="mask",
|
||||
confidence=1.0,
|
||||
crop_region=CropRegion(0, 0, 2, 2),
|
||||
bbox=BoundingBox(0, 0, 2, 2),
|
||||
label=label,
|
||||
)
|
||||
@@ -0,0 +1,89 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Batch SEGS Comfy v3 wrapper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment
|
||||
from simple_syrup.nodes_v3.batch_segs import BatchSEGSV3
|
||||
|
||||
|
||||
def test_batch_segs_v3_schema_uses_autogrow_segs_inputs() -> None:
|
||||
"""The v3 schema exposes expandable SEGS inputs."""
|
||||
|
||||
schema = BatchSEGSV3.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.BatchSEGS"
|
||||
assert schema.display_name == "Batch SEGS"
|
||||
assert schema.category == "SimpleSyrup/Detection"
|
||||
assert [input_item.id for input_item in schema.inputs] == ["segs_inputs"]
|
||||
assert schema.inputs[0].io_type == "COMFY_AUTOGROW_V3"
|
||||
assert schema.inputs[0].template.prefix == "segs"
|
||||
assert schema.inputs[0].template.min == 2
|
||||
assert schema.inputs[0].template.max == 50
|
||||
assert schema.inputs[0].template.input.io_type == "SEGS"
|
||||
assert [output.id for output in schema.outputs] == ["segs"]
|
||||
assert schema.outputs[0].io_type == "SEGS"
|
||||
|
||||
|
||||
def test_batch_segs_v3_execute_returns_impact_compatible_segs() -> None:
|
||||
"""The v3 wrapper batches SEGS in Autogrow insertion order."""
|
||||
|
||||
first = (
|
||||
(16, 16),
|
||||
[
|
||||
_segment("1", CropRegion(0, 0, 2, 2)),
|
||||
_segment("2", CropRegion(2, 0, 4, 2)),
|
||||
_segment("3", CropRegion(4, 0, 6, 2)),
|
||||
],
|
||||
)
|
||||
second = (
|
||||
(16, 16),
|
||||
(
|
||||
_segment("4", CropRegion(0, 2, 2, 4)),
|
||||
_segment("5", CropRegion(2, 2, 4, 4)),
|
||||
_segment("6", CropRegion(4, 2, 6, 4)),
|
||||
),
|
||||
)
|
||||
|
||||
(raw_segs,) = BatchSEGSV3.execute({"segs0": first, "segs1": second})
|
||||
segs = cast(tuple[tuple[int, int], list[Segment]], raw_segs)
|
||||
|
||||
header, segments = segs
|
||||
assert header == (16, 16)
|
||||
assert isinstance(segments, list)
|
||||
assert [segment.label for segment in segments] == ["1", "2", "3", "4", "5", "6"]
|
||||
|
||||
|
||||
def test_batch_segs_v3_execute_surfaces_header_mismatch() -> None:
|
||||
"""The v3 wrapper keeps domain validation errors visible."""
|
||||
|
||||
first = ((8, 16), (_segment("first", CropRegion(0, 0, 2, 2)),))
|
||||
second = ((16, 8), (_segment("second", CropRegion(0, 0, 2, 2)),))
|
||||
|
||||
with pytest.raises(ValueError, match="input 2 is 16x8 but input 1 is 8x16"):
|
||||
BatchSEGSV3.execute({"segs0": first, "segs1": second})
|
||||
|
||||
|
||||
def _segment(label: str, crop_region: CropRegion) -> Segment:
|
||||
"""Create a segment for Batch SEGS v3 tests."""
|
||||
|
||||
return Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask="mask",
|
||||
confidence=1.0,
|
||||
crop_region=crop_region,
|
||||
bbox=BoundingBox(
|
||||
crop_region.left,
|
||||
crop_region.top,
|
||||
crop_region.right,
|
||||
crop_region.bottom,
|
||||
),
|
||||
label=label,
|
||||
)
|
||||
@@ -10,6 +10,7 @@ import pytest
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import (
|
||||
ConditioningBatch,
|
||||
batch_conditioning,
|
||||
select_conditioning,
|
||||
split_prompt_batch,
|
||||
)
|
||||
@@ -65,6 +66,25 @@ def test_conditioning_batch_rejects_negative_indexes() -> None:
|
||||
ConditioningBatch(("a",)).select(-1)
|
||||
|
||||
|
||||
def test_batch_conditioning_flattens_batches_and_normal_conditioning() -> None:
|
||||
"""Mixed conditioning inputs become one ordered per-region batch."""
|
||||
|
||||
first = ConditioningBatch(("auto 1", "auto 2"))
|
||||
hand = "hand 1"
|
||||
second = ConditioningBatch(("auto 3",))
|
||||
|
||||
batch = batch_conditioning((first, hand, second))
|
||||
|
||||
assert batch.entries == ("auto 1", "auto 2", "hand 1", "auto 3")
|
||||
|
||||
|
||||
def test_batch_conditioning_rejects_no_inputs() -> None:
|
||||
"""At least one input is needed to build a conditioning batch."""
|
||||
|
||||
with pytest.raises(ValueError, match="one or more inputs"):
|
||||
batch_conditioning(())
|
||||
|
||||
|
||||
def test_select_conditioning_broadcasts_normal_conditioning() -> None:
|
||||
"""Normal conditionings pass through unchanged for any valid index."""
|
||||
|
||||
|
||||
@@ -297,8 +297,8 @@ def test_noise_mask_false_omits_latent_mask_but_composites_pixels() -> None:
|
||||
assert torch.all(result.image[:, 4:, :, :] == 0.0)
|
||||
|
||||
|
||||
def test_noise_mask_feather_applies_differential_diffusion() -> None:
|
||||
"""Feathered denoise masks request differential diffusion patching."""
|
||||
def test_noise_mask_feather_requests_single_clone_differential_diffusion() -> None:
|
||||
"""Feathered denoise masks are composed inside the regional runtime clone."""
|
||||
|
||||
sampler = _FakeRegionalSampler()
|
||||
service = _service(sampler)
|
||||
@@ -314,8 +314,9 @@ def test_noise_mask_feather_applies_differential_diffusion() -> None:
|
||||
**(_settings() | {"noise_mask": True, "noise_mask_feather": 2}),
|
||||
)
|
||||
|
||||
assert sampler.patch_count == 1
|
||||
assert sampler.sample_calls[0]["model"] == "patched model"
|
||||
assert sampler.patch_count == 0
|
||||
assert sampler.sample_calls[0]["model"] == "model"
|
||||
assert sampler.sample_calls[0]["differential_diffusion"] is True
|
||||
|
||||
|
||||
def test_encode_decode_and_sampler_controls_are_forwarded() -> None:
|
||||
@@ -479,6 +480,7 @@ class _FakeRegionalSampler:
|
||||
denoise: float,
|
||||
global_prompt_weight: float,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Record regional sample options and return the latent unchanged."""
|
||||
|
||||
@@ -497,6 +499,7 @@ class _FakeRegionalSampler:
|
||||
"denoise": denoise,
|
||||
"global_prompt_weight": global_prompt_weight,
|
||||
"preview_context": preview_context,
|
||||
"differential_diffusion": differential_diffusion,
|
||||
}
|
||||
)
|
||||
return latent_image
|
||||
|
||||
@@ -211,6 +211,29 @@ def test_tiled_noise_mask_feather_keeps_sampled_crop_geometry() -> None:
|
||||
assert float(noise_mask[0, 1, 1]) > float(noise_mask[0, 0, 1])
|
||||
|
||||
|
||||
def test_tiled_noise_mask_feather_requests_single_clone_differential_diffusion() -> (
|
||||
None
|
||||
):
|
||||
"""Feathered denoise masks are composed inside the tiled runtime clone."""
|
||||
|
||||
sampler = _FakeTiledSampler()
|
||||
service = _service(sampler)
|
||||
|
||||
service.detail(
|
||||
_image(),
|
||||
_segs(_segment()),
|
||||
"model",
|
||||
"vae",
|
||||
[],
|
||||
[],
|
||||
**(_settings() | {"noise_mask": True, "noise_mask_feather": 2}),
|
||||
)
|
||||
|
||||
assert sampler.patch_count == 0
|
||||
assert sampler.sample_calls[0]["model"] == "model"
|
||||
assert sampler.sample_calls[0]["differential_diffusion"] is True
|
||||
|
||||
|
||||
def test_tiled_detail_sampler_delegates_to_shared_sampling_service() -> None:
|
||||
"""The detailer adapter uses the shared tiled diffusion dispatch service."""
|
||||
|
||||
@@ -241,6 +264,7 @@ def test_tiled_detail_sampler_delegates_to_shared_sampling_service() -> None:
|
||||
latent_tile_overlap=12,
|
||||
latent_tile_batch_size=3,
|
||||
preview_context=preview_context,
|
||||
differential_diffusion=True,
|
||||
)
|
||||
|
||||
assert result is latent
|
||||
@@ -248,6 +272,7 @@ def test_tiled_detail_sampler_delegates_to_shared_sampling_service() -> None:
|
||||
assert call["diffusion_mode"] == "mixture_of_diffusers"
|
||||
assert call["latent_image"] is latent
|
||||
assert call["preview_context"] is preview_context
|
||||
assert call["differential_diffusion"] is True
|
||||
|
||||
|
||||
class _FakeTiledSamplingService:
|
||||
@@ -277,6 +302,7 @@ class _FakeTiledSamplingService:
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Record tiled sampling arguments and return the latent unchanged."""
|
||||
|
||||
@@ -298,6 +324,7 @@ class _FakeTiledSamplingService:
|
||||
"latent_tile_overlap": latent_tile_overlap,
|
||||
"latent_tile_batch_size": latent_tile_batch_size,
|
||||
"preview_context": preview_context,
|
||||
"differential_diffusion": differential_diffusion,
|
||||
}
|
||||
)
|
||||
return latent_image
|
||||
@@ -356,6 +383,7 @@ class _FakeTiledSampler:
|
||||
latent_tile_overlap: int,
|
||||
latent_tile_batch_size: int,
|
||||
preview_context: DetailPreviewContext | None = None,
|
||||
differential_diffusion: bool = False,
|
||||
) -> Latent:
|
||||
"""Record tiled sample options and return the latent unchanged."""
|
||||
|
||||
@@ -377,6 +405,7 @@ class _FakeTiledSampler:
|
||||
"latent_tile_overlap": latent_tile_overlap,
|
||||
"latent_tile_batch_size": latent_tile_batch_size,
|
||||
"preview_context": preview_context,
|
||||
"differential_diffusion": differential_diffusion,
|
||||
}
|
||||
)
|
||||
return latent_image
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the OpenAI-compatible external LLM HTTP client."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.domain.external_llm import (
|
||||
ExternalLLMChatRequest,
|
||||
ExternalLLMProviderError,
|
||||
)
|
||||
from simple_syrup.runtime.external_llm_client import ExternalLLMClient
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
"""Context-managed HTTP response double."""
|
||||
|
||||
status = 200
|
||||
|
||||
def __init__(self, payload: object) -> None:
|
||||
"""Store JSON payload bytes."""
|
||||
|
||||
self._payload = json.dumps(payload).encode("utf-8")
|
||||
|
||||
def read(self) -> bytes:
|
||||
"""Return response bytes."""
|
||||
|
||||
return self._payload
|
||||
|
||||
def __enter__(self) -> FakeResponse:
|
||||
"""Enter context manager."""
|
||||
|
||||
return self
|
||||
|
||||
def __exit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc: BaseException | None,
|
||||
traceback: object | None,
|
||||
) -> None:
|
||||
"""Exit context manager."""
|
||||
|
||||
|
||||
class CapturingUrlopen:
|
||||
"""Transport double that captures requests."""
|
||||
|
||||
def __init__(self, payload: object) -> None:
|
||||
"""Store the response payload."""
|
||||
|
||||
self.payload = payload
|
||||
self.requests: list[tuple[urllib.request.Request, int]] = []
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
request: urllib.request.Request,
|
||||
timeout: int,
|
||||
) -> FakeResponse:
|
||||
"""Capture request and return configured response."""
|
||||
|
||||
self.requests.append((request, timeout))
|
||||
return FakeResponse(self.payload)
|
||||
|
||||
|
||||
def test_list_models_sends_authenticated_models_request() -> None:
|
||||
"""Model refresh uses the OpenAI-compatible `/models` endpoint."""
|
||||
|
||||
urlopen = CapturingUrlopen(
|
||||
{"data": [{"id": "model-a"}, {"id": "model-a"}, {"id": "model-b"}]}
|
||||
)
|
||||
client = ExternalLLMClient(urlopen)
|
||||
|
||||
models = client.list_models("https://provider.example/v1/", "secret")
|
||||
|
||||
request, timeout = urlopen.requests[0]
|
||||
assert request.full_url == "https://provider.example/v1/models"
|
||||
assert request.get_method() == "GET"
|
||||
assert request.headers["Authorization"] == "Bearer secret"
|
||||
assert timeout == 10
|
||||
assert models == ("model-a", "model-b")
|
||||
|
||||
|
||||
def test_list_models_rejects_malformed_response() -> None:
|
||||
"""Malformed provider model payloads fail clearly."""
|
||||
|
||||
client = ExternalLLMClient(CapturingUrlopen({"data": "bad"}))
|
||||
|
||||
with pytest.raises(ExternalLLMProviderError, match="models response"):
|
||||
client.list_models("https://provider.example/v1", "secret")
|
||||
|
||||
|
||||
def test_chat_completion_sends_messages_and_returns_content() -> None:
|
||||
"""Chat completion uses system and user messages and extracts content."""
|
||||
|
||||
urlopen = CapturingUrlopen(
|
||||
{"choices": [{"message": {"content": "rewritten prompt"}}]}
|
||||
)
|
||||
client = ExternalLLMClient(urlopen)
|
||||
|
||||
response = client.create_chat_completion(
|
||||
"https://provider.example/v1",
|
||||
"secret",
|
||||
ExternalLLMChatRequest(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
),
|
||||
)
|
||||
|
||||
request, timeout = urlopen.requests[0]
|
||||
assert request.full_url == "https://provider.example/v1/chat/completions"
|
||||
assert request.get_method() == "POST"
|
||||
assert timeout == 60
|
||||
assert isinstance(request.data, bytes)
|
||||
assert json.loads(request.data.decode("utf-8")) == {
|
||||
"model": "model-a",
|
||||
"messages": [
|
||||
{"role": "system", "content": "system"},
|
||||
{"role": "user", "content": "user"},
|
||||
],
|
||||
"max_tokens": 1024,
|
||||
}
|
||||
assert response.content == "rewritten prompt"
|
||||
|
||||
|
||||
def test_chat_completion_sends_configured_max_tokens() -> None:
|
||||
"""Chat completion forwards the workflow response token limit."""
|
||||
|
||||
urlopen = CapturingUrlopen({"choices": [{"message": {"content": "ok"}}]})
|
||||
client = ExternalLLMClient(urlopen)
|
||||
|
||||
client.create_chat_completion(
|
||||
"https://provider.example/v1",
|
||||
"secret",
|
||||
ExternalLLMChatRequest(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
max_tokens=128,
|
||||
),
|
||||
)
|
||||
|
||||
request, _timeout = urlopen.requests[0]
|
||||
assert isinstance(request.data, bytes)
|
||||
body = json.loads(request.data.decode("utf-8"))
|
||||
assert body["max_tokens"] == 128
|
||||
|
||||
|
||||
def test_chat_completion_sends_reasoning_effort_when_requested() -> None:
|
||||
"""Reasoning effort high/medium/low is sent as a top-level field."""
|
||||
|
||||
urlopen = CapturingUrlopen({"choices": [{"message": {"content": "ok"}}]})
|
||||
client = ExternalLLMClient(urlopen)
|
||||
|
||||
client.create_chat_completion(
|
||||
"https://provider.example/v1",
|
||||
"secret",
|
||||
ExternalLLMChatRequest(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
reasoning_effort="low",
|
||||
),
|
||||
)
|
||||
|
||||
request, _timeout = urlopen.requests[0]
|
||||
assert isinstance(request.data, bytes)
|
||||
body = json.loads(request.data.decode("utf-8"))
|
||||
assert body["reasoning_effort"] == "low"
|
||||
assert "chat_template_kwargs" not in body
|
||||
|
||||
|
||||
def test_chat_completion_sends_thinking_disabled_for_off() -> None:
|
||||
"""The off option maps to chat template thinking disablement."""
|
||||
|
||||
urlopen = CapturingUrlopen({"choices": [{"message": {"content": "ok"}}]})
|
||||
client = ExternalLLMClient(urlopen)
|
||||
|
||||
client.create_chat_completion(
|
||||
"https://provider.example/v1",
|
||||
"secret",
|
||||
ExternalLLMChatRequest(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
reasoning_effort="off",
|
||||
),
|
||||
)
|
||||
|
||||
request, _timeout = urlopen.requests[0]
|
||||
assert isinstance(request.data, bytes)
|
||||
body = json.loads(request.data.decode("utf-8"))
|
||||
assert body["chat_template_kwargs"] == {"thinking": False}
|
||||
assert "reasoning_effort" not in body
|
||||
|
||||
|
||||
def test_chat_completion_sends_image_content_when_image_is_present() -> None:
|
||||
"""Vision requests use OpenAI-compatible multimodal user content."""
|
||||
|
||||
urlopen = CapturingUrlopen({"choices": [{"message": {"content": "caption"}}]})
|
||||
client = ExternalLLMClient(urlopen)
|
||||
|
||||
client.create_chat_completion(
|
||||
"https://provider.example/v1",
|
||||
"secret",
|
||||
ExternalLLMChatRequest(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="describe this",
|
||||
image_data_url="data:image/png;base64,abc",
|
||||
),
|
||||
)
|
||||
|
||||
request, _timeout = urlopen.requests[0]
|
||||
assert isinstance(request.data, bytes)
|
||||
body = json.loads(request.data.decode("utf-8"))
|
||||
assert body["messages"][1] == {
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "describe this"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,abc"},
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def test_chat_completion_rejects_malformed_response() -> None:
|
||||
"""Missing assistant content is reported as a provider response error."""
|
||||
|
||||
client = ExternalLLMClient(CapturingUrlopen({"choices": []}))
|
||||
|
||||
with pytest.raises(ExternalLLMProviderError, match="chat completion"):
|
||||
client.create_chat_completion(
|
||||
"https://provider.example/v1",
|
||||
"secret",
|
||||
ExternalLLMChatRequest(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_transport_failures_become_provider_errors() -> None:
|
||||
"""URL transport errors are wrapped as provider errors."""
|
||||
|
||||
def failing_urlopen(
|
||||
_request: urllib.request.Request,
|
||||
_timeout: int,
|
||||
) -> FakeResponse:
|
||||
raise urllib.error.URLError("offline")
|
||||
|
||||
client = ExternalLLMClient(failing_urlopen)
|
||||
|
||||
with pytest.raises(ExternalLLMProviderError, match="request failed"):
|
||||
client.list_models("https://provider.example/v1", "secret")
|
||||
@@ -0,0 +1,25 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for external LLM image encoding."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.runtime.external_llm_images import ExternalLLMImageEncoder
|
||||
|
||||
|
||||
def test_image_encoder_returns_png_data_url() -> None:
|
||||
"""ComfyUI IMAGE tensors are encoded as PNG data URLs."""
|
||||
|
||||
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
|
||||
|
||||
data_url = ExternalLLMImageEncoder().encode_first_image_as_data_url(image)
|
||||
|
||||
assert data_url.startswith("data:image/png;base64,")
|
||||
payload = data_url.removeprefix("data:image/png;base64,")
|
||||
assert base64.b64decode(payload).startswith(b"\x89PNG")
|
||||
@@ -0,0 +1,98 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for external LLM keyring credential storage."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.domain.external_llm import ExternalLLMConfigError
|
||||
from simple_syrup.runtime.external_llm_keyring import (
|
||||
ExternalLLMKeyringError,
|
||||
ExternalLLMKeyringStore,
|
||||
credential_username,
|
||||
)
|
||||
|
||||
|
||||
class FakeKeyring:
|
||||
"""In-memory keyring double."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create empty credential storage."""
|
||||
|
||||
self.passwords: dict[tuple[str, str], str] = {}
|
||||
|
||||
def get_password(self, service_name: str, username: str) -> str | None:
|
||||
"""Return a stored password."""
|
||||
|
||||
return self.passwords.get((service_name, username))
|
||||
|
||||
def set_password(self, service_name: str, username: str, password: str) -> None:
|
||||
"""Store a password."""
|
||||
|
||||
self.passwords[(service_name, username)] = password
|
||||
|
||||
def delete_password(self, service_name: str, username: str) -> None:
|
||||
"""Delete a password."""
|
||||
|
||||
self.passwords.pop((service_name, username), None)
|
||||
|
||||
|
||||
class FailingKeyring(FakeKeyring):
|
||||
"""Keyring double that fails all operations."""
|
||||
|
||||
def get_password(self, service_name: str, username: str) -> str | None:
|
||||
"""Raise on read."""
|
||||
|
||||
raise RuntimeError("backend failed")
|
||||
|
||||
def set_password(self, service_name: str, username: str, password: str) -> None:
|
||||
"""Raise on write."""
|
||||
|
||||
raise RuntimeError("backend failed")
|
||||
|
||||
def delete_password(self, service_name: str, username: str) -> None:
|
||||
"""Raise on delete."""
|
||||
|
||||
raise RuntimeError("backend failed")
|
||||
|
||||
|
||||
def test_keyring_store_saves_reads_checks_and_deletes_api_key() -> None:
|
||||
"""API keys are stored under the normalized endpoint credential name."""
|
||||
|
||||
keyring = FakeKeyring()
|
||||
store = ExternalLLMKeyringStore(keyring)
|
||||
|
||||
store.save_api_key("https://provider.example/v1/", " secret ")
|
||||
|
||||
assert store.has_api_key("https://provider.example/v1") is True
|
||||
assert store.get_api_key("https://provider.example/v1") == "secret"
|
||||
assert credential_username("https://provider.example/v1/").endswith(
|
||||
"https://provider.example/v1"
|
||||
)
|
||||
|
||||
store.delete_api_key("https://provider.example/v1")
|
||||
|
||||
assert store.has_api_key("https://provider.example/v1") is False
|
||||
|
||||
|
||||
def test_keyring_store_rejects_empty_api_key() -> None:
|
||||
"""Empty API keys fail before touching credential storage."""
|
||||
|
||||
store = ExternalLLMKeyringStore(FakeKeyring())
|
||||
|
||||
with pytest.raises(ExternalLLMConfigError, match="must not be empty"):
|
||||
store.save_api_key("https://provider.example/v1", " ")
|
||||
|
||||
|
||||
def test_keyring_failures_do_not_expose_api_key() -> None:
|
||||
"""Credential backend failures are wrapped without leaking secrets."""
|
||||
|
||||
store = ExternalLLMKeyringStore(FailingKeyring())
|
||||
|
||||
with pytest.raises(ExternalLLMKeyringError) as error:
|
||||
store.save_api_key("https://provider.example/v1", "secret-value")
|
||||
|
||||
assert "secret-value" not in str(error.value)
|
||||
@@ -0,0 +1,102 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the legacy External LLM Prompt node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from simple_syrup.nodes.external_llm_prompt import ExternalLLMPrompt
|
||||
from simple_syrup.services.external_llm_prompt_service import CONFIGURE_EXTERNAL_LLM
|
||||
|
||||
|
||||
class FakeExternalLLMService:
|
||||
"""Service double for node tests."""
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return deterministic model choices."""
|
||||
|
||||
return ["model-a"]
|
||||
|
||||
def generate(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = 1024,
|
||||
reasoning_effort: str = "default",
|
||||
image: object | None = None,
|
||||
) -> str:
|
||||
"""Return a deterministic response."""
|
||||
|
||||
return (
|
||||
f"{model}:{system_prompt}:{user_prompt}:"
|
||||
f"{max_tokens}:{reasoning_effort}:{image}"
|
||||
)
|
||||
|
||||
|
||||
def test_external_llm_prompt_input_types_are_cached_and_single_line() -> None:
|
||||
"""The node exposes cached model choices and single-line prompt widgets."""
|
||||
|
||||
original_service = ExternalLLMPrompt._service
|
||||
ExternalLLMPrompt._service = FakeExternalLLMService() # type: ignore[assignment]
|
||||
try:
|
||||
input_types = ExternalLLMPrompt.INPUT_TYPES()
|
||||
finally:
|
||||
ExternalLLMPrompt._service = original_service
|
||||
|
||||
inputs = input_types["required"]
|
||||
assert inputs["model"][0] == ["model-a"]
|
||||
assert inputs["model"][1]["default"] == "model-a"
|
||||
assert inputs["system_prompt"][0] == "STRING"
|
||||
assert inputs["system_prompt"][1]["multiline"] is False
|
||||
assert inputs["user_prompt"][0] == "STRING"
|
||||
assert inputs["user_prompt"][1]["multiline"] is False
|
||||
assert inputs["max_tokens"][0] == "INT"
|
||||
assert inputs["max_tokens"][1]["default"] == 1024
|
||||
assert inputs["max_tokens"][1]["min"] == 1
|
||||
assert inputs["max_tokens"][1]["max"] == 32768
|
||||
assert inputs["reasoning_effort"][0] == [
|
||||
"default",
|
||||
"high",
|
||||
"medium",
|
||||
"low",
|
||||
"off",
|
||||
]
|
||||
assert inputs["reasoning_effort"][1]["default"] == "default"
|
||||
assert input_types["optional"]["image"][0] == "IMAGE"
|
||||
|
||||
|
||||
def test_external_llm_prompt_output_contract() -> None:
|
||||
"""The node returns one named STRING output."""
|
||||
|
||||
assert ExternalLLMPrompt.RETURN_TYPES == ("STRING",)
|
||||
assert ExternalLLMPrompt.RETURN_NAMES == ("response",)
|
||||
assert ExternalLLMPrompt.FUNCTION == "generate"
|
||||
|
||||
|
||||
def test_external_llm_prompt_executes_service() -> None:
|
||||
"""Node execution delegates to the prompt service."""
|
||||
|
||||
node = ExternalLLMPrompt()
|
||||
original_service = node._service
|
||||
node._service = FakeExternalLLMService() # type: ignore[assignment]
|
||||
try:
|
||||
result = node.generate(
|
||||
"model-a",
|
||||
"system",
|
||||
"user",
|
||||
128,
|
||||
"off",
|
||||
image="image",
|
||||
)
|
||||
finally:
|
||||
node._service = original_service
|
||||
|
||||
assert result == ("model-a:system:user:128:off:image",)
|
||||
|
||||
|
||||
def test_external_llm_prompt_sentinel_constant_matches_plan() -> None:
|
||||
"""The setup sentinel is stable for dropdown fallback behavior."""
|
||||
|
||||
assert CONFIGURE_EXTERNAL_LLM == "Configure external LLM endpoint"
|
||||
@@ -0,0 +1,441 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for external LLM prompt service orchestration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.domain.external_llm import (
|
||||
ExternalLLMChatRequest,
|
||||
ExternalLLMChatResponse,
|
||||
ExternalLLMConfigError,
|
||||
ExternalLLMProviderError,
|
||||
)
|
||||
from simple_syrup.runtime.settings import ExternalLLMSettings, SimpleSyrupSettings
|
||||
from simple_syrup.services.external_llm_prompt_service import (
|
||||
CONFIGURE_EXTERNAL_LLM,
|
||||
ExternalLLMPromptService,
|
||||
)
|
||||
|
||||
|
||||
class FakeSettingsRepository:
|
||||
"""In-memory settings repository."""
|
||||
|
||||
def __init__(self, settings: SimpleSyrupSettings) -> None:
|
||||
"""Store initial settings."""
|
||||
|
||||
self.settings = settings
|
||||
|
||||
def load(self) -> SimpleSyrupSettings:
|
||||
"""Return current settings."""
|
||||
|
||||
return self.settings
|
||||
|
||||
def save(self, settings: SimpleSyrupSettings) -> SimpleSyrupSettings:
|
||||
"""Persist current settings."""
|
||||
|
||||
self.settings = settings
|
||||
return settings
|
||||
|
||||
|
||||
class FakeKeyStore:
|
||||
"""Credential store double."""
|
||||
|
||||
def __init__(self, api_key: str = "") -> None:
|
||||
"""Store a fake API key."""
|
||||
|
||||
self.api_key = api_key
|
||||
|
||||
def has_api_key(self, _base_url: str) -> bool:
|
||||
"""Return whether a key exists."""
|
||||
|
||||
return bool(self.api_key)
|
||||
|
||||
def get_api_key(self, _base_url: str) -> str:
|
||||
"""Return the fake key."""
|
||||
|
||||
return self.api_key
|
||||
|
||||
def save_api_key(self, _base_url: str, api_key: str) -> None:
|
||||
"""Save the fake key."""
|
||||
|
||||
self.api_key = api_key
|
||||
|
||||
def delete_api_key(self, _base_url: str) -> None:
|
||||
"""Delete the fake key."""
|
||||
|
||||
self.api_key = ""
|
||||
|
||||
|
||||
class FakeClient:
|
||||
"""Provider client double."""
|
||||
|
||||
def __init__(self, models: tuple[str, ...] = ("model-a",)) -> None:
|
||||
"""Store provider models."""
|
||||
|
||||
self.models = models
|
||||
self.requests: list[ExternalLLMChatRequest] = []
|
||||
self.list_model_calls = 0
|
||||
|
||||
def list_models(self, _base_url: str, _api_key: str) -> tuple[str, ...]:
|
||||
"""Return configured models."""
|
||||
|
||||
self.list_model_calls += 1
|
||||
return self.models
|
||||
|
||||
def create_chat_completion(
|
||||
self,
|
||||
_base_url: str,
|
||||
_api_key: str,
|
||||
request: ExternalLLMChatRequest,
|
||||
) -> ExternalLLMChatResponse:
|
||||
"""Capture request and return content."""
|
||||
|
||||
self.requests.append(request)
|
||||
return ExternalLLMChatResponse("assistant response")
|
||||
|
||||
|
||||
class FailingModelClient(FakeClient):
|
||||
"""Provider client double that cannot refresh models."""
|
||||
|
||||
def list_models(self, _base_url: str, _api_key: str) -> tuple[str, ...]:
|
||||
"""Raise a provider failure during model refresh."""
|
||||
|
||||
self.list_model_calls += 1
|
||||
raise ExternalLLMProviderError("The external LLM provider rejected /models.")
|
||||
|
||||
|
||||
class FakeImageEncoder:
|
||||
"""Image encoder double."""
|
||||
|
||||
def encode_first_image_as_data_url(self, image: object) -> str:
|
||||
"""Return a deterministic data URL."""
|
||||
|
||||
assert image == "image"
|
||||
return "data:image/png;base64,abc"
|
||||
|
||||
|
||||
def test_model_choices_return_sentinel_without_cached_models() -> None:
|
||||
"""No cached models produces a clear sentinel choice."""
|
||||
|
||||
service = ExternalLLMPromptService(
|
||||
FakeSettingsRepository(SimpleSyrupSettings()),
|
||||
FakeKeyStore(),
|
||||
FakeClient(),
|
||||
)
|
||||
|
||||
assert service.model_choices() == [CONFIGURE_EXTERNAL_LLM]
|
||||
|
||||
|
||||
def test_refresh_models_skips_when_unconfigured() -> None:
|
||||
"""Refresh does no network work without endpoint credentials."""
|
||||
|
||||
repository = FakeSettingsRepository(SimpleSyrupSettings())
|
||||
client = FakeClient(("model-a",))
|
||||
service = ExternalLLMPromptService(repository, FakeKeyStore(), client)
|
||||
|
||||
assert service.refresh_models() == ()
|
||||
assert repository.load().external_llm.cached_models == ()
|
||||
|
||||
|
||||
def test_refresh_models_persists_provider_models() -> None:
|
||||
"""Refresh stores sanitized provider model choices."""
|
||||
|
||||
repository = FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("old-model",),
|
||||
default_model="old-model",
|
||||
)
|
||||
)
|
||||
)
|
||||
service = ExternalLLMPromptService(
|
||||
repository,
|
||||
FakeKeyStore("secret"),
|
||||
FakeClient(("model-a", "model-b")),
|
||||
)
|
||||
|
||||
assert service.refresh_models() == ("model-a", "model-b")
|
||||
assert repository.load().external_llm.default_model == "model-a"
|
||||
|
||||
|
||||
def test_save_config_clears_cached_models_when_endpoint_changes() -> None:
|
||||
"""Provider model cache belongs to the configured endpoint."""
|
||||
|
||||
repository = FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://old-provider.example/v1",
|
||||
cached_models=("old-model",),
|
||||
default_model="old-model",
|
||||
)
|
||||
)
|
||||
)
|
||||
service = ExternalLLMPromptService(repository, FakeKeyStore(), FakeClient())
|
||||
|
||||
saved = service.save_config("https://new-provider.example/v1")
|
||||
|
||||
assert saved.external_llm.base_url == "https://new-provider.example/v1"
|
||||
assert saved.external_llm.cached_models == ()
|
||||
assert saved.external_llm.default_model == ""
|
||||
|
||||
|
||||
def test_save_config_preserves_cached_models_for_same_endpoint() -> None:
|
||||
"""Saving the existing endpoint keeps its model cache and default choice."""
|
||||
|
||||
repository = FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("model-a", "model-b"),
|
||||
default_model="model-a",
|
||||
)
|
||||
)
|
||||
)
|
||||
service = ExternalLLMPromptService(repository, FakeKeyStore(), FakeClient())
|
||||
|
||||
saved = service.save_config(
|
||||
"https://provider.example/v1/",
|
||||
default_model="model-b",
|
||||
)
|
||||
|
||||
assert saved.external_llm.cached_models == ("model-a", "model-b")
|
||||
assert saved.external_llm.default_model == "model-b"
|
||||
|
||||
|
||||
def test_save_config_refreshes_models_when_endpoint_has_key() -> None:
|
||||
"""Saving an endpoint updates the available model cache when possible."""
|
||||
|
||||
repository = FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("old-model",),
|
||||
default_model="old-model",
|
||||
)
|
||||
)
|
||||
)
|
||||
client = FakeClient(("model-a", "model-b"))
|
||||
service = ExternalLLMPromptService(repository, FakeKeyStore("secret"), client)
|
||||
|
||||
saved = service.save_config("https://provider.example/v1")
|
||||
|
||||
assert client.list_model_calls == 1
|
||||
assert saved.external_llm.cached_models == ("model-a", "model-b")
|
||||
assert saved.external_llm.default_model == "model-a"
|
||||
|
||||
|
||||
def test_save_config_succeeds_when_model_refresh_fails() -> None:
|
||||
"""Endpoint storage is not blocked by provider model discovery failures."""
|
||||
|
||||
repository = FakeSettingsRepository(SimpleSyrupSettings())
|
||||
service = ExternalLLMPromptService(
|
||||
repository,
|
||||
FakeKeyStore("secret"),
|
||||
FailingModelClient(),
|
||||
)
|
||||
|
||||
saved = service.save_config("https://provider.example/v1")
|
||||
|
||||
assert saved.external_llm.base_url == "https://provider.example/v1"
|
||||
assert saved.external_llm.cached_models == ()
|
||||
|
||||
|
||||
def test_save_api_key_persists_key_when_model_refresh_fails() -> None:
|
||||
"""Credential storage is independent from provider model discovery."""
|
||||
|
||||
key_store = FakeKeyStore()
|
||||
service = ExternalLLMPromptService(
|
||||
FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
)
|
||||
)
|
||||
),
|
||||
key_store,
|
||||
FailingModelClient(),
|
||||
)
|
||||
|
||||
saved = service.save_api_key("secret")
|
||||
|
||||
assert key_store.has_api_key(saved.external_llm.base_url) is True
|
||||
assert saved.external_llm.cached_models == ()
|
||||
|
||||
|
||||
def test_generate_rejects_sentinel_when_unconfigured() -> None:
|
||||
"""The setup sentinel still requires endpoint configuration."""
|
||||
|
||||
service = ExternalLLMPromptService(
|
||||
FakeSettingsRepository(SimpleSyrupSettings()),
|
||||
FakeKeyStore(),
|
||||
FakeClient(),
|
||||
)
|
||||
|
||||
with pytest.raises(ExternalLLMConfigError, match="Configure an external LLM"):
|
||||
service.generate(CONFIGURE_EXTERNAL_LLM, "", "prompt")
|
||||
|
||||
|
||||
def test_generate_refreshes_stale_sentinel_and_uses_default_model() -> None:
|
||||
"""Execution recovers when the graph still holds the setup sentinel."""
|
||||
|
||||
repository = FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
)
|
||||
)
|
||||
)
|
||||
client = FakeClient(("model-a", "model-b"))
|
||||
service = ExternalLLMPromptService(repository, FakeKeyStore("secret"), client)
|
||||
|
||||
assert service.generate(CONFIGURE_EXTERNAL_LLM, "system", "user") == (
|
||||
"assistant response"
|
||||
)
|
||||
assert client.list_model_calls == 1
|
||||
assert client.requests[0].model == "model-a"
|
||||
assert repository.load().external_llm.cached_models == ("model-a", "model-b")
|
||||
|
||||
|
||||
def test_generate_uses_cached_default_when_sentinel_is_stale() -> None:
|
||||
"""A stale sentinel resolves to the saved default without provider refresh."""
|
||||
|
||||
repository = FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("model-a", "model-b"),
|
||||
default_model="model-b",
|
||||
)
|
||||
)
|
||||
)
|
||||
client = FakeClient(("model-c",))
|
||||
service = ExternalLLMPromptService(repository, FakeKeyStore("secret"), client)
|
||||
|
||||
assert service.generate(CONFIGURE_EXTERNAL_LLM, "system", "user") == (
|
||||
"assistant response"
|
||||
)
|
||||
assert client.list_model_calls == 0
|
||||
assert client.requests[0].model == "model-b"
|
||||
|
||||
|
||||
def test_generate_rejects_empty_user_prompt() -> None:
|
||||
"""User prompt validation happens before provider work."""
|
||||
|
||||
service = configured_service()
|
||||
|
||||
with pytest.raises(ExternalLLMConfigError, match="User prompt"):
|
||||
service.generate("model-a", "", " ")
|
||||
|
||||
|
||||
def test_generate_returns_provider_response() -> None:
|
||||
"""Prompt generation delegates to the provider client."""
|
||||
|
||||
service = configured_service()
|
||||
|
||||
assert service.generate("model-a", "system", "user") == "assistant response"
|
||||
|
||||
|
||||
def test_generate_forwards_max_tokens() -> None:
|
||||
"""Prompt generation includes the requested response token limit."""
|
||||
|
||||
client = FakeClient()
|
||||
service = configured_service(client=client)
|
||||
|
||||
assert service.generate("model-a", "system", "user", max_tokens=128) == (
|
||||
"assistant response"
|
||||
)
|
||||
assert client.requests[0].max_tokens == 128
|
||||
|
||||
|
||||
def test_generate_forwards_reasoning_effort() -> None:
|
||||
"""Prompt generation includes the requested reasoning behavior."""
|
||||
|
||||
client = FakeClient()
|
||||
service = configured_service(client=client)
|
||||
|
||||
assert (
|
||||
service.generate(
|
||||
"model-a",
|
||||
"system",
|
||||
"user",
|
||||
reasoning_effort="off",
|
||||
)
|
||||
== "assistant response"
|
||||
)
|
||||
assert client.requests[0].reasoning_effort == "off"
|
||||
|
||||
|
||||
def test_generate_rejects_invalid_max_tokens() -> None:
|
||||
"""Response token limit must be a positive integer."""
|
||||
|
||||
service = configured_service()
|
||||
|
||||
with pytest.raises(ExternalLLMConfigError, match="max tokens"):
|
||||
service.generate("model-a", "system", "user", max_tokens=0)
|
||||
|
||||
|
||||
def test_generate_rejects_invalid_reasoning_effort() -> None:
|
||||
"""Reasoning effort must be one of the node choices."""
|
||||
|
||||
service = configured_service()
|
||||
|
||||
with pytest.raises(ExternalLLMConfigError, match="reasoning effort"):
|
||||
service.generate("model-a", "system", "user", reasoning_effort="none")
|
||||
|
||||
|
||||
def test_generate_attaches_encoded_image_when_supplied() -> None:
|
||||
"""Optional image inputs are encoded before provider chat requests."""
|
||||
|
||||
client = FakeClient()
|
||||
service = configured_service(client=client, image_encoder=FakeImageEncoder())
|
||||
|
||||
assert (
|
||||
service.generate("model-a", "system", "user", image="image")
|
||||
== "assistant response"
|
||||
)
|
||||
assert client.requests[0].image_data_url == "data:image/png;base64,abc"
|
||||
|
||||
|
||||
def test_generate_with_image_data_url_forwards_preencoded_image() -> None:
|
||||
"""SEG callers can supply a prebuilt image data URL."""
|
||||
|
||||
client = FakeClient()
|
||||
service = configured_service(client=client)
|
||||
|
||||
assert (
|
||||
service.generate_with_image_data_url(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
image_data_url="data:image/png;base64,seg",
|
||||
)
|
||||
== "assistant response"
|
||||
)
|
||||
assert client.requests[0].image_data_url == "data:image/png;base64,seg"
|
||||
|
||||
|
||||
def configured_service(
|
||||
client: FakeClient | None = None,
|
||||
image_encoder: FakeImageEncoder | None = None,
|
||||
) -> ExternalLLMPromptService:
|
||||
"""Create a service with endpoint, cached model, key, and fake client."""
|
||||
|
||||
return ExternalLLMPromptService(
|
||||
FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("model-a",),
|
||||
default_model="model-a",
|
||||
)
|
||||
)
|
||||
),
|
||||
FakeKeyStore("secret"),
|
||||
client or FakeClient(),
|
||||
image_encoder=image_encoder,
|
||||
)
|
||||
@@ -0,0 +1,81 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the External LLM Prompt Comfy v3 wrapper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from simple_syrup.nodes.external_llm_prompt import ExternalLLMPrompt
|
||||
from simple_syrup.nodes_v3.external_llm_prompt import ExternalLLMPromptV3
|
||||
|
||||
|
||||
class FakeExternalLLMService:
|
||||
"""Service double for v3 wrapper tests."""
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return deterministic model choices."""
|
||||
|
||||
return ["model-a"]
|
||||
|
||||
def generate(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = 1024,
|
||||
reasoning_effort: str = "default",
|
||||
image: object | None = None,
|
||||
) -> str:
|
||||
"""Return a deterministic response."""
|
||||
|
||||
return (
|
||||
f"{model}:{system_prompt}:{user_prompt}:"
|
||||
f"{max_tokens}:{reasoning_effort}:{image}"
|
||||
)
|
||||
|
||||
|
||||
def test_external_llm_prompt_v3_schema_includes_max_tokens() -> None:
|
||||
"""The v3 schema exposes the response token limit."""
|
||||
|
||||
schema = ExternalLLMPromptV3.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.ExternalLLMPrompt"
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"model",
|
||||
"system_prompt",
|
||||
"user_prompt",
|
||||
"max_tokens",
|
||||
"reasoning_effort",
|
||||
"image",
|
||||
]
|
||||
max_tokens = schema.inputs[3]
|
||||
assert max_tokens.io_type == "INT"
|
||||
assert max_tokens.default == 1024
|
||||
assert max_tokens.min == 1
|
||||
assert max_tokens.max == 32768
|
||||
reasoning_effort = schema.inputs[4]
|
||||
assert reasoning_effort.io_type == "COMBO"
|
||||
assert reasoning_effort.options == ["default", "high", "medium", "low", "off"]
|
||||
assert reasoning_effort.default == "default"
|
||||
|
||||
|
||||
def test_external_llm_prompt_v3_execute_forwards_max_tokens(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""The v3 wrapper forwards max_tokens to the legacy implementation."""
|
||||
|
||||
monkeypatch.setattr(ExternalLLMPrompt, "_service", FakeExternalLLMService())
|
||||
|
||||
result = ExternalLLMPromptV3.execute(
|
||||
model="model-a",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
max_tokens=128,
|
||||
reasoning_effort="off",
|
||||
image="image",
|
||||
)
|
||||
|
||||
assert result == ("model-a:system:user:128:off:image",)
|
||||
@@ -0,0 +1,401 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for external LLM settings HTTP routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from typing import cast
|
||||
|
||||
from aiohttp import web
|
||||
|
||||
from simple_syrup.domain.external_llm import (
|
||||
ExternalLLMChatRequest,
|
||||
ExternalLLMChatResponse,
|
||||
ExternalLLMProviderError,
|
||||
)
|
||||
from simple_syrup.runtime.external_llm_routes import (
|
||||
EXTERNAL_LLM_API_KEY_ROUTE,
|
||||
EXTERNAL_LLM_MODELS_REFRESH_ROUTE,
|
||||
EXTERNAL_LLM_SETTINGS_ROUTE,
|
||||
Handler,
|
||||
PromptServerProtocol,
|
||||
register_external_llm_routes,
|
||||
)
|
||||
from simple_syrup.runtime.settings import ExternalLLMSettings, SimpleSyrupSettings
|
||||
from simple_syrup.services.external_llm_prompt_service import ExternalLLMPromptService
|
||||
|
||||
|
||||
class FakeRoutes:
|
||||
"""Route table double that records handlers."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create an empty route table."""
|
||||
|
||||
self.get_handlers: dict[str, Handler] = {}
|
||||
self.post_handlers: dict[str, Handler] = {}
|
||||
self.delete_handlers: dict[str, Handler] = {}
|
||||
|
||||
def get(self, path: str) -> Callable[[Handler], Handler]:
|
||||
"""Record GET handlers."""
|
||||
|
||||
def decorator(handler: Handler) -> Handler:
|
||||
self.get_handlers[path] = handler
|
||||
return handler
|
||||
|
||||
return decorator
|
||||
|
||||
def post(self, path: str) -> Callable[[Handler], Handler]:
|
||||
"""Record POST handlers."""
|
||||
|
||||
def decorator(handler: Handler) -> Handler:
|
||||
self.post_handlers[path] = handler
|
||||
return handler
|
||||
|
||||
return decorator
|
||||
|
||||
def delete(self, path: str) -> Callable[[Handler], Handler]:
|
||||
"""Record DELETE handlers."""
|
||||
|
||||
def decorator(handler: Handler) -> Handler:
|
||||
self.delete_handlers[path] = handler
|
||||
return handler
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
class FakePromptServer:
|
||||
"""PromptServer double exposing route table."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create fake prompt server."""
|
||||
|
||||
self.routes = FakeRoutes()
|
||||
|
||||
|
||||
class FakeRequest:
|
||||
"""Request double with injectable JSON payload."""
|
||||
|
||||
def __init__(self, payload: object) -> None:
|
||||
"""Store payload."""
|
||||
|
||||
self._payload = payload
|
||||
|
||||
async def json(self) -> object:
|
||||
"""Return configured payload."""
|
||||
|
||||
return self._payload
|
||||
|
||||
|
||||
class FakeSettingsRepository:
|
||||
"""In-memory settings repository."""
|
||||
|
||||
def __init__(self, settings: SimpleSyrupSettings) -> None:
|
||||
"""Store initial settings."""
|
||||
|
||||
self.settings = settings
|
||||
|
||||
def load(self) -> SimpleSyrupSettings:
|
||||
"""Return current settings."""
|
||||
|
||||
return self.settings
|
||||
|
||||
def save(self, settings: SimpleSyrupSettings) -> SimpleSyrupSettings:
|
||||
"""Persist settings."""
|
||||
|
||||
self.settings = settings
|
||||
return settings
|
||||
|
||||
|
||||
class FakeKeyStore:
|
||||
"""Credential store double."""
|
||||
|
||||
def __init__(self, api_key: str = "") -> None:
|
||||
"""Store key."""
|
||||
|
||||
self.api_key = api_key
|
||||
|
||||
def has_api_key(self, _base_url: str) -> bool:
|
||||
"""Return whether a key exists."""
|
||||
|
||||
return bool(self.api_key)
|
||||
|
||||
def get_api_key(self, _base_url: str) -> str:
|
||||
"""Return key."""
|
||||
|
||||
return self.api_key
|
||||
|
||||
def save_api_key(self, _base_url: str, api_key: str) -> None:
|
||||
"""Save key."""
|
||||
|
||||
self.api_key = api_key
|
||||
|
||||
def delete_api_key(self, _base_url: str) -> None:
|
||||
"""Delete key."""
|
||||
|
||||
self.api_key = ""
|
||||
|
||||
|
||||
class FakeClient:
|
||||
"""Provider client double."""
|
||||
|
||||
def list_models(self, _base_url: str, _api_key: str) -> tuple[str, ...]:
|
||||
"""Return fake models."""
|
||||
|
||||
return ("model-a", "model-b")
|
||||
|
||||
def create_chat_completion(
|
||||
self,
|
||||
_base_url: str,
|
||||
_api_key: str,
|
||||
_request: ExternalLLMChatRequest,
|
||||
) -> ExternalLLMChatResponse:
|
||||
"""Return fake content."""
|
||||
|
||||
return ExternalLLMChatResponse("response")
|
||||
|
||||
|
||||
class FailingModelClient(FakeClient):
|
||||
"""Provider client double that fails model discovery."""
|
||||
|
||||
def list_models(self, _base_url: str, _api_key: str) -> tuple[str, ...]:
|
||||
"""Raise a provider error during model refresh."""
|
||||
|
||||
raise ExternalLLMProviderError("The external LLM provider rejected /models.")
|
||||
|
||||
|
||||
def test_external_llm_route_registration_records_handlers() -> None:
|
||||
"""External LLM routes register with Comfy's route table."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
|
||||
assert register_fake_routes(configured_service(), prompt_server) is True
|
||||
|
||||
assert EXTERNAL_LLM_SETTINGS_ROUTE in prompt_server.routes.get_handlers
|
||||
assert EXTERNAL_LLM_SETTINGS_ROUTE in prompt_server.routes.post_handlers
|
||||
assert EXTERNAL_LLM_API_KEY_ROUTE in prompt_server.routes.post_handlers
|
||||
assert EXTERNAL_LLM_API_KEY_ROUTE in prompt_server.routes.delete_handlers
|
||||
assert EXTERNAL_LLM_MODELS_REFRESH_ROUTE in prompt_server.routes.post_handlers
|
||||
|
||||
|
||||
def test_get_settings_returns_non_secret_payload() -> None:
|
||||
"""GET returns endpoint settings and key presence without the key."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
register_fake_routes(configured_service(api_key="secret"), prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.get_handlers[EXTERNAL_LLM_SETTINGS_ROUTE](object())
|
||||
)
|
||||
|
||||
payload = json.loads(response_text(response))
|
||||
assert response.status == 200
|
||||
assert payload == {
|
||||
"base_url": "https://provider.example/v1",
|
||||
"cached_models": ["model-a"],
|
||||
"default_model": "model-a",
|
||||
"has_api_key": True,
|
||||
}
|
||||
assert "secret" not in response_text(response)
|
||||
|
||||
|
||||
def test_post_settings_validates_and_saves_payload() -> None:
|
||||
"""POST settings persists non-secret endpoint config."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
service = ExternalLLMPromptService(
|
||||
FakeSettingsRepository(SimpleSyrupSettings()),
|
||||
FakeKeyStore(),
|
||||
FakeClient(),
|
||||
)
|
||||
register_fake_routes(service, prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.post_handlers[EXTERNAL_LLM_SETTINGS_ROUTE](
|
||||
FakeRequest(
|
||||
{
|
||||
"base_url": "https://provider.example/v1/",
|
||||
"default_model": "",
|
||||
}
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
payload = json.loads(response_text(response))
|
||||
assert response.status == 200
|
||||
assert payload["base_url"] == "https://provider.example/v1"
|
||||
|
||||
|
||||
def test_post_settings_refreshes_models_when_key_exists() -> None:
|
||||
"""Saving endpoint settings updates model choices when credentials exist."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
service = ExternalLLMPromptService(
|
||||
FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("old-model",),
|
||||
default_model="old-model",
|
||||
)
|
||||
)
|
||||
),
|
||||
FakeKeyStore("secret"),
|
||||
FakeClient(),
|
||||
)
|
||||
register_fake_routes(service, prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.post_handlers[EXTERNAL_LLM_SETTINGS_ROUTE](
|
||||
FakeRequest(
|
||||
{
|
||||
"base_url": "https://provider.example/v1",
|
||||
"default_model": "",
|
||||
}
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
payload = json.loads(response_text(response))
|
||||
assert response.status == 200
|
||||
assert payload["cached_models"] == ["model-a", "model-b"]
|
||||
|
||||
|
||||
def test_post_api_key_stores_key_and_refreshes_models() -> None:
|
||||
"""API key route stores the key without returning it."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
register_fake_routes(configured_service(), prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.post_handlers[EXTERNAL_LLM_API_KEY_ROUTE](
|
||||
FakeRequest({"api_key": "secret"})
|
||||
)
|
||||
)
|
||||
|
||||
payload = json.loads(response_text(response))
|
||||
assert response.status == 200
|
||||
assert payload["has_api_key"] is True
|
||||
assert payload["cached_models"] == ["model-a", "model-b"]
|
||||
assert "secret" not in response_text(response)
|
||||
|
||||
|
||||
def test_post_api_key_succeeds_when_model_refresh_fails() -> None:
|
||||
"""API key storage still succeeds if immediate model discovery fails."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
service = ExternalLLMPromptService(
|
||||
FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
)
|
||||
)
|
||||
),
|
||||
FakeKeyStore(),
|
||||
FailingModelClient(),
|
||||
)
|
||||
register_fake_routes(service, prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.post_handlers[EXTERNAL_LLM_API_KEY_ROUTE](
|
||||
FakeRequest({"api_key": "secret"})
|
||||
)
|
||||
)
|
||||
|
||||
payload = json.loads(response_text(response))
|
||||
assert response.status == 200
|
||||
assert payload["has_api_key"] is True
|
||||
assert payload["cached_models"] == []
|
||||
assert "secret" not in response_text(response)
|
||||
|
||||
|
||||
def test_delete_api_key_removes_key_presence() -> None:
|
||||
"""DELETE removes the stored API key."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
register_fake_routes(configured_service(api_key="secret"), prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.delete_handlers[EXTERNAL_LLM_API_KEY_ROUTE](object())
|
||||
)
|
||||
|
||||
assert response.status == 200
|
||||
assert json.loads(response_text(response))["has_api_key"] is False
|
||||
|
||||
|
||||
def test_refresh_models_updates_cache_when_configured() -> None:
|
||||
"""Refresh route updates cached model ids."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
register_fake_routes(configured_service(api_key="secret"), prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.post_handlers[EXTERNAL_LLM_MODELS_REFRESH_ROUTE](object())
|
||||
)
|
||||
|
||||
assert response.status == 200
|
||||
assert json.loads(response_text(response))["cached_models"] == [
|
||||
"model-a",
|
||||
"model-b",
|
||||
]
|
||||
|
||||
|
||||
def test_refresh_models_skips_when_unconfigured() -> None:
|
||||
"""Refresh route succeeds without endpoint credentials."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
service = ExternalLLMPromptService(
|
||||
FakeSettingsRepository(SimpleSyrupSettings()),
|
||||
FakeKeyStore(),
|
||||
FakeClient(),
|
||||
)
|
||||
register_fake_routes(service, prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.post_handlers[EXTERNAL_LLM_MODELS_REFRESH_ROUTE](object())
|
||||
)
|
||||
|
||||
assert response.status == 200
|
||||
assert json.loads(response_text(response))["cached_models"] == []
|
||||
|
||||
|
||||
def configured_service(api_key: str = "") -> ExternalLLMPromptService:
|
||||
"""Create a configured route service."""
|
||||
|
||||
return ExternalLLMPromptService(
|
||||
FakeSettingsRepository(
|
||||
SimpleSyrupSettings(
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("model-a",),
|
||||
default_model="model-a",
|
||||
)
|
||||
)
|
||||
),
|
||||
FakeKeyStore(api_key),
|
||||
FakeClient(),
|
||||
)
|
||||
|
||||
|
||||
def register_fake_routes(
|
||||
service: ExternalLLMPromptService,
|
||||
prompt_server: FakePromptServer,
|
||||
) -> bool:
|
||||
"""Register routes against a fake server."""
|
||||
|
||||
return register_external_llm_routes(
|
||||
service,
|
||||
cast(PromptServerProtocol, prompt_server),
|
||||
)
|
||||
|
||||
|
||||
def response_text(response: web.Response) -> str:
|
||||
"""Return response text after asserting it exists."""
|
||||
|
||||
assert response.text is not None
|
||||
return response.text
|
||||
@@ -0,0 +1,160 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for external LLM SEG crop image encoding."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from io import BytesIO
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment
|
||||
from simple_syrup.runtime.external_llm_images import ExternalLLMSegsImageEncoder
|
||||
|
||||
|
||||
def test_segs_image_encoder_returns_transparent_mask_png() -> None:
|
||||
"""Transparent mode hides outside-mask pixels with PNG alpha."""
|
||||
|
||||
segment = _segment(torch.tensor([[1.0, 0.0], [0.0, 1.0]]))
|
||||
|
||||
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
|
||||
_image(),
|
||||
segment,
|
||||
"transparent mask",
|
||||
)
|
||||
|
||||
image = _decode_png(encoded)
|
||||
assert image.mode == "RGBA"
|
||||
assert image.size == (2, 2)
|
||||
assert _rgba_pixel(image, 0, 0)[3] == 255
|
||||
assert _rgba_pixel(image, 1, 0)[3] == 0
|
||||
assert _rgba_pixel(image, 0, 1)[3] == 0
|
||||
assert _rgba_pixel(image, 1, 1)[3] == 255
|
||||
|
||||
|
||||
def test_segs_image_encoder_returns_black_mask_png() -> None:
|
||||
"""Black mode zeros outside-mask pixels and keeps RGB output."""
|
||||
|
||||
segment = _segment(torch.tensor([[1.0, 0.0], [0.0, 1.0]]))
|
||||
|
||||
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
|
||||
_image(),
|
||||
segment,
|
||||
"black mask",
|
||||
)
|
||||
|
||||
image = _decode_png(encoded)
|
||||
assert image.mode == "RGB"
|
||||
assert image.size == (2, 2)
|
||||
assert image.getpixel((1, 0)) == (0, 0, 0)
|
||||
assert image.getpixel((0, 1)) == (0, 0, 0)
|
||||
assert image.getpixel((0, 0)) != (0, 0, 0)
|
||||
assert image.getpixel((1, 1)) != (0, 0, 0)
|
||||
|
||||
|
||||
def test_segs_image_encoder_returns_full_crop_png() -> None:
|
||||
"""Full crop mode preserves the whole crop rectangle."""
|
||||
|
||||
segment = _segment(torch.zeros((2, 2)))
|
||||
|
||||
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
|
||||
_image(),
|
||||
segment,
|
||||
"full crop",
|
||||
)
|
||||
|
||||
image = _decode_png(encoded)
|
||||
assert image.mode == "RGB"
|
||||
assert image.size == (2, 2)
|
||||
assert image.getpixel((0, 0)) != (0, 0, 0)
|
||||
assert image.getpixel((1, 0)) != (0, 0, 0)
|
||||
assert image.getpixel((0, 1)) != (0, 0, 0)
|
||||
assert image.getpixel((1, 1)) != (0, 0, 0)
|
||||
|
||||
|
||||
def test_segs_image_encoder_accepts_full_image_masks() -> None:
|
||||
"""Full-image SEG masks are cropped to the SEG crop region."""
|
||||
|
||||
full_mask = torch.zeros((4, 4), dtype=torch.float32)
|
||||
full_mask[1, 1] = 1.0
|
||||
full_mask[2, 2] = 1.0
|
||||
segment = _segment(full_mask)
|
||||
|
||||
encoded = ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
|
||||
_image(),
|
||||
segment,
|
||||
"transparent mask",
|
||||
)
|
||||
|
||||
image = _decode_png(encoded)
|
||||
assert [_rgba_pixel(image, x, y)[3] for y in range(2) for x in range(2)] == [
|
||||
255,
|
||||
0,
|
||||
0,
|
||||
255,
|
||||
]
|
||||
|
||||
|
||||
def test_segs_image_encoder_rejects_invalid_masks() -> None:
|
||||
"""SEG masks must be HW or BHW tensors."""
|
||||
|
||||
segment = _segment(torch.zeros((1, 1, 1, 1), dtype=torch.float32))
|
||||
|
||||
with pytest.raises(ValueError, match="HW or BHW"):
|
||||
ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
|
||||
_image(),
|
||||
segment,
|
||||
"transparent mask",
|
||||
)
|
||||
|
||||
|
||||
def test_segs_image_encoder_rejects_unknown_mode() -> None:
|
||||
"""SEG image mode is restricted to the node combo choices."""
|
||||
|
||||
with pytest.raises(ValueError, match="seg_image_mode"):
|
||||
ExternalLLMSegsImageEncoder().encode_segment_as_data_url(
|
||||
_image(),
|
||||
_segment(torch.ones((2, 2), dtype=torch.float32)),
|
||||
"white mask",
|
||||
)
|
||||
|
||||
|
||||
def _decode_png(data_url: str) -> Image.Image:
|
||||
"""Decode a PNG data URL into a PIL image."""
|
||||
|
||||
assert data_url.startswith("data:image/png;base64,")
|
||||
payload = data_url.removeprefix("data:image/png;base64,")
|
||||
return Image.open(BytesIO(base64.b64decode(payload)))
|
||||
|
||||
|
||||
def _rgba_pixel(image: Image.Image, x: int, y: int) -> tuple[int, int, int, int]:
|
||||
"""Return one RGBA pixel with a precise test type."""
|
||||
|
||||
return cast(tuple[int, int, int, int], image.getpixel((x, y)))
|
||||
|
||||
|
||||
def _segment(mask: torch.Tensor) -> Segment:
|
||||
"""Return a SEG that crops a stable two-by-two image region."""
|
||||
|
||||
return Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask=mask,
|
||||
confidence=1.0,
|
||||
crop_region=CropRegion(1, 1, 3, 3),
|
||||
bbox=BoundingBox(1, 1, 3, 3),
|
||||
label="seg",
|
||||
)
|
||||
|
||||
|
||||
def _image() -> torch.Tensor:
|
||||
"""Return a deterministic BHWC image with no black crop pixels."""
|
||||
|
||||
return (
|
||||
torch.arange(1, 4 * 4 * 3 + 1, dtype=torch.float32).reshape(1, 4, 4, 3) / 255.0
|
||||
)
|
||||
@@ -47,6 +47,33 @@ def test_direct_vae_decode_resolves_samples_and_vae_links() -> None:
|
||||
assert result.vae_link == ("loader", 2)
|
||||
|
||||
|
||||
def test_vae_decode_options_resolves_samples_and_vae_links() -> None:
|
||||
"""VAE Decode (Options) output resolves to the source latent and VAE."""
|
||||
|
||||
prompt = {
|
||||
"decode": {
|
||||
"class_type": "SimpleSyrup.VAEDecodeOptions",
|
||||
"inputs": {
|
||||
"use_tiling": True,
|
||||
"samples": ["latent", 0],
|
||||
"vae": ["loader", 2],
|
||||
"tile_size": 512,
|
||||
"overlap": 64,
|
||||
"temporal_size": 64,
|
||||
"temporal_overlap": 8,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
result = trace_vae_decode_provenance(prompt, ["decode", 0], {})
|
||||
|
||||
assert isinstance(result, VaeDecodeProvenance)
|
||||
assert result.decode_node_id == "decode"
|
||||
assert result.image_output == ("decode", 0)
|
||||
assert result.samples_link == ("latent", 0)
|
||||
assert result.vae_link == ("loader", 2)
|
||||
|
||||
|
||||
def test_transparent_node_resolves_to_upstream_vae_decode() -> None:
|
||||
"""A declared pass-through node is traversed to its source link."""
|
||||
|
||||
@@ -137,6 +164,25 @@ def test_vae_decode_non_image_output_breaks_provenance() -> None:
|
||||
assert result.reason == "VAEDecode output is not the image output"
|
||||
|
||||
|
||||
def test_vae_decode_options_non_image_output_breaks_provenance() -> None:
|
||||
"""Only VAE Decode (Options) output slot 0 is trusted as image provenance."""
|
||||
|
||||
result = trace_vae_decode_provenance(
|
||||
{
|
||||
"decode": {
|
||||
"class_type": "SimpleSyrup.VAEDecodeOptions",
|
||||
"inputs": {"samples": ["latent", 0], "vae": ["loader", 2]},
|
||||
}
|
||||
},
|
||||
["decode", 1],
|
||||
{},
|
||||
)
|
||||
|
||||
assert isinstance(result, BrokenProvenance)
|
||||
assert result.reason == "VAEDecode output is not the image output"
|
||||
assert result.class_type == "SimpleSyrup.VAEDecodeOptions"
|
||||
|
||||
|
||||
def test_missing_samples_link_breaks_provenance() -> None:
|
||||
"""VAEDecode without a graph-linked samples input cannot supply provenance."""
|
||||
|
||||
@@ -150,6 +196,25 @@ def test_missing_samples_link_breaks_provenance() -> None:
|
||||
assert result.reason == "VAEDecode samples input is not a graph link"
|
||||
|
||||
|
||||
def test_vae_decode_options_missing_samples_link_breaks_provenance() -> None:
|
||||
"""VAE Decode (Options) requires a graph-linked samples input."""
|
||||
|
||||
result = trace_vae_decode_provenance(
|
||||
{
|
||||
"decode": {
|
||||
"class_type": "SimpleSyrup.VAEDecodeOptions",
|
||||
"inputs": {"samples": "latent"},
|
||||
}
|
||||
},
|
||||
["decode", 0],
|
||||
{},
|
||||
)
|
||||
|
||||
assert isinstance(result, BrokenProvenance)
|
||||
assert result.reason == "VAEDecode samples input is not a graph link"
|
||||
assert result.class_type == "SimpleSyrup.VAEDecodeOptions"
|
||||
|
||||
|
||||
def test_missing_node_breaks_provenance() -> None:
|
||||
"""Missing source nodes produce broken provenance."""
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.nodes.ksampler_extras import KSamplerExtras
|
||||
from simple_syrup.runtime import sampling_samplers, sampling_schedulers
|
||||
|
||||
@@ -63,6 +64,8 @@ def test_input_types_match_simple_ksampler_contract() -> None:
|
||||
"latent_image",
|
||||
"denoise",
|
||||
)
|
||||
assert required["positive"][0] == "CONDITIONING,CONDITIONING_BATCH"
|
||||
assert required["negative"][0] == "CONDITIONING,CONDITIONING_BATCH"
|
||||
|
||||
|
||||
def test_node_metadata_matches_contract() -> None:
|
||||
@@ -289,3 +292,107 @@ def test_sample_delegates_to_runtime_helpers(
|
||||
assert calls["sample_custom"]["sampler"] is sampler
|
||||
assert calls["sample_custom"]["sigmas"] is fixed_sigmas
|
||||
assert calls["sample_custom"]["disable_pbar"] is True
|
||||
|
||||
|
||||
def test_sample_selects_conditioning_batch_per_latent_item(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""Conditioning batches are selected before calling Comfy sampling."""
|
||||
|
||||
calls: list[dict[str, Any]] = []
|
||||
model = FakeModel()
|
||||
sampler = FakeSampler()
|
||||
latent_samples = torch.arange(2 * 4 * 2 * 2, dtype=torch.float32).reshape(
|
||||
(2, 4, 2, 2)
|
||||
)
|
||||
fixed_noise = torch.full_like(latent_samples, 2.0)
|
||||
fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32)
|
||||
noise_mask = torch.ones((2, 1, 2, 2), dtype=torch.float32)
|
||||
latent_image: dict[str, Any] = {
|
||||
"samples": latent_samples,
|
||||
"batch_index": [4, 9],
|
||||
"noise_mask": noise_mask,
|
||||
}
|
||||
|
||||
monkeypatch.setattr(
|
||||
sampling_samplers,
|
||||
"resolve_sampler",
|
||||
lambda sampler_name: sampler,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
sampling_schedulers,
|
||||
"calculate_sigmas",
|
||||
lambda **kwargs: fixed_sigmas,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
comfy_sample,
|
||||
"fix_empty_latent_channels",
|
||||
lambda model, samples, downscale_ratio_spacial: samples,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
comfy_sample,
|
||||
"prepare_noise",
|
||||
lambda samples, seed, batch_inds: fixed_noise,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
latent_preview,
|
||||
"prepare_callback",
|
||||
lambda received_model, steps: "callback",
|
||||
)
|
||||
|
||||
def fake_sample_custom(
|
||||
received_model: FakeModel,
|
||||
noise: torch.Tensor,
|
||||
cfg: float,
|
||||
received_sampler: FakeSampler,
|
||||
sigmas: torch.Tensor,
|
||||
positive: object,
|
||||
negative: object,
|
||||
latent_image: torch.Tensor,
|
||||
noise_mask: torch.Tensor | None,
|
||||
callback: str,
|
||||
disable_pbar: bool,
|
||||
seed: int,
|
||||
) -> torch.Tensor:
|
||||
"""Record one per-item sample call and return a marked tensor."""
|
||||
|
||||
del received_model, cfg, received_sampler, sigmas, callback, disable_pbar, seed
|
||||
calls.append(
|
||||
{
|
||||
"noise": noise,
|
||||
"positive": positive,
|
||||
"negative": negative,
|
||||
"latent_image": latent_image,
|
||||
"noise_mask": noise_mask,
|
||||
}
|
||||
)
|
||||
return torch.full_like(latent_image, float(len(calls)))
|
||||
|
||||
monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom)
|
||||
monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False)
|
||||
|
||||
(output,) = KSamplerExtras().sample(
|
||||
model=model,
|
||||
seed=123,
|
||||
steps=2,
|
||||
cfg=7.5,
|
||||
sampler_name="lcm",
|
||||
scheduler="GITS",
|
||||
positive=ConditioningBatch(("positive-0", "positive-1")),
|
||||
negative=ConditioningBatch(("negative-last",)),
|
||||
latent_image=latent_image,
|
||||
denoise=0.8,
|
||||
)
|
||||
|
||||
assert len(calls) == 2
|
||||
assert calls[0]["positive"] == "positive-0"
|
||||
assert calls[1]["positive"] == "positive-1"
|
||||
assert calls[0]["negative"] == "negative-last"
|
||||
assert calls[1]["negative"] == "negative-last"
|
||||
assert torch.equal(calls[0]["noise"], fixed_noise[0:1])
|
||||
assert torch.equal(calls[1]["noise"], fixed_noise[1:2])
|
||||
assert torch.equal(calls[0]["noise_mask"], noise_mask[0:1])
|
||||
assert torch.equal(calls[1]["noise_mask"], noise_mask[1:2])
|
||||
assert output["samples"].shape == latent_samples.shape
|
||||
assert torch.equal(output["samples"][0], torch.full((4, 2, 2), 1.0))
|
||||
assert torch.equal(output["samples"][1], torch.full((4, 2, 2), 2.0))
|
||||
|
||||
@@ -50,6 +50,8 @@ def test_input_types_match_tiled_diffusion_contract(
|
||||
"mixture_of_diffusers",
|
||||
]
|
||||
assert required["diffusion_mode"][1]["default"] == "multidiffusion"
|
||||
assert required["positive"][0] == "CONDITIONING,CONDITIONING_BATCH"
|
||||
assert required["negative"][0] == "CONDITIONING,CONDITIONING_BATCH"
|
||||
assert required["latent_tile_width"][1]["default"] == 128
|
||||
assert required["latent_tile_width"][1]["max"] == 512
|
||||
assert required["latent_tile_height"][1]["default"] == 128
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for v3 wrappers that delegate to existing implementation classes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from simple_syrup.nodes_v3.legacy_node_wrappers import LegacyNodeV3Adapter
|
||||
|
||||
|
||||
class _FakeHidden:
|
||||
"""Hidden holder double for adapter execution tests."""
|
||||
|
||||
prompt = {"node": "metadata"}
|
||||
|
||||
|
||||
class _FakeLegacyNode:
|
||||
"""Legacy implementation double with visible and hidden inputs."""
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("result",)
|
||||
OUTPUT_TOOLTIPS = ("Delegated result.",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "SimpleSyrup/Test"
|
||||
DESCRIPTION = "Delegates test inputs."
|
||||
SEARCH_ALIASES = ["adapter"]
|
||||
|
||||
calls: ClassVar[list[tuple[str, object]]] = []
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...] | str]]:
|
||||
"""Return a small legacy contract with one hidden input."""
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"text": (
|
||||
"STRING",
|
||||
{"default": "hello", "tooltip": "Text to delegate."},
|
||||
),
|
||||
},
|
||||
"hidden": {"prompt": "PROMPT"},
|
||||
}
|
||||
|
||||
def run(self, text: str, prompt: object | None = None) -> tuple[str]:
|
||||
"""Record delegated inputs and return the visible value."""
|
||||
|
||||
self.calls.append((text, prompt))
|
||||
return (text,)
|
||||
|
||||
|
||||
class _FakeAdapter(LegacyNodeV3Adapter):
|
||||
"""Concrete adapter used to verify generic wrapper behavior."""
|
||||
|
||||
LEGACY_NODE_CLASS = _FakeLegacyNode
|
||||
NODE_ID = "SimpleSyrup.FakeAdapter"
|
||||
DISPLAY_NAME = "Fake Adapter"
|
||||
hidden = _FakeHidden()
|
||||
|
||||
|
||||
def test_legacy_node_v3_adapter_builds_schema() -> None:
|
||||
"""The adapter converts legacy metadata into a v3 schema."""
|
||||
|
||||
schema = _FakeAdapter.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.FakeAdapter"
|
||||
assert schema.display_name == "Fake Adapter"
|
||||
assert schema.category == "SimpleSyrup/Test"
|
||||
assert schema.inputs[0].id == "text"
|
||||
assert schema.inputs[0].tooltip == "Text to delegate."
|
||||
assert schema.hidden[0].value == "PROMPT"
|
||||
assert schema.outputs[0].id == "result"
|
||||
assert schema.outputs[0].tooltip == "Delegated result."
|
||||
|
||||
|
||||
def test_legacy_node_v3_adapter_delegates_execution_with_hidden_inputs() -> None:
|
||||
"""The adapter passes v3 visible and hidden values to the implementation."""
|
||||
|
||||
_FakeLegacyNode.calls.clear()
|
||||
|
||||
assert _FakeAdapter.execute(text="value") == ("value",)
|
||||
assert _FakeLegacyNode.calls == [("value", {"node": "metadata"})]
|
||||
@@ -83,6 +83,7 @@ class FakeModel:
|
||||
def __init__(
|
||||
self,
|
||||
model_options: dict[str, Any] | None = None,
|
||||
parent: FakeModel | None = None,
|
||||
) -> None:
|
||||
"""Create a fake model patcher."""
|
||||
|
||||
@@ -90,11 +91,14 @@ class FakeModel:
|
||||
self.model_options = {} if model_options is None else model_options
|
||||
self.wrapper: Any = None
|
||||
self.model_sampling = object()
|
||||
self.parent = parent
|
||||
self.clone_count = 0
|
||||
|
||||
def clone(self) -> FakeModel:
|
||||
"""Return a cloned model with copied options."""
|
||||
|
||||
return FakeModel(self.model_options.copy())
|
||||
self.clone_count += 1
|
||||
return FakeModel(self.model_options.copy(), parent=self)
|
||||
|
||||
def set_model_unet_function_wrapper(self, wrapper: object) -> None:
|
||||
"""Capture the installed model function wrapper."""
|
||||
@@ -102,6 +106,11 @@ class FakeModel:
|
||||
self.wrapper = wrapper
|
||||
self.model_options["model_function_wrapper"] = wrapper
|
||||
|
||||
def set_model_denoise_mask_function(self, denoise_mask_function: object) -> None:
|
||||
"""Capture the installed denoise-mask function."""
|
||||
|
||||
self.model_options["denoise_mask_function"] = denoise_mask_function
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Return the requested fake model object."""
|
||||
|
||||
@@ -119,6 +128,31 @@ class FakeSampler:
|
||||
return None
|
||||
|
||||
|
||||
def test_clone_model_with_mixture_composes_differential_on_same_clone() -> None:
|
||||
"""Differential diffusion is installed without cloning a temporary parent."""
|
||||
|
||||
model = FakeModel()
|
||||
|
||||
wrapped_model, _plan = mod_sampling.clone_model_with_mixture_of_diffusers(
|
||||
model,
|
||||
latent_width=8,
|
||||
latent_height=4,
|
||||
tile_width=4,
|
||||
tile_height=4,
|
||||
overlap=0,
|
||||
tile_batch_size=2,
|
||||
differential_diffusion=True,
|
||||
)
|
||||
|
||||
assert model.clone_count == 1
|
||||
assert wrapped_model.parent is model
|
||||
assert callable(wrapped_model.model_options["denoise_mask_function"])
|
||||
assert isinstance(
|
||||
wrapped_model.wrapper,
|
||||
mod_sampling.MixtureOfDiffusersModelWrapper,
|
||||
)
|
||||
|
||||
|
||||
def test_model_wrapper_tiles_input_conditioning_and_transformer_options() -> None:
|
||||
"""The wrapper tiles latents, conditioning tensors, timesteps, and metadata."""
|
||||
|
||||
|
||||
@@ -91,6 +91,7 @@ class FakeModel:
|
||||
def __init__(
|
||||
self,
|
||||
model_options: dict[str, Any] | None = None,
|
||||
parent: FakeModel | None = None,
|
||||
) -> None:
|
||||
"""Create a fake model patcher."""
|
||||
|
||||
@@ -98,11 +99,14 @@ class FakeModel:
|
||||
self.model_options = {} if model_options is None else model_options
|
||||
self.wrapper: Any = None
|
||||
self.model_sampling = object()
|
||||
self.parent = parent
|
||||
self.clone_count = 0
|
||||
|
||||
def clone(self) -> FakeModel:
|
||||
"""Return a cloned model with copied options."""
|
||||
|
||||
return FakeModel(self.model_options.copy())
|
||||
self.clone_count += 1
|
||||
return FakeModel(self.model_options.copy(), parent=self)
|
||||
|
||||
def set_model_unet_function_wrapper(self, wrapper: object) -> None:
|
||||
"""Capture the installed model function wrapper."""
|
||||
@@ -110,6 +114,11 @@ class FakeModel:
|
||||
self.wrapper = wrapper
|
||||
self.model_options["model_function_wrapper"] = wrapper
|
||||
|
||||
def set_model_denoise_mask_function(self, denoise_mask_function: object) -> None:
|
||||
"""Capture the installed denoise-mask function."""
|
||||
|
||||
self.model_options["denoise_mask_function"] = denoise_mask_function
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Return the requested fake model object."""
|
||||
|
||||
@@ -150,6 +159,31 @@ def test_clone_model_with_multidiffusion_installs_wrapper() -> None:
|
||||
assert plan.tile_batch_size == 2
|
||||
|
||||
|
||||
def test_clone_model_with_multidiffusion_composes_differential_on_same_clone() -> None:
|
||||
"""Differential diffusion is installed without cloning a temporary parent."""
|
||||
|
||||
model = FakeModel()
|
||||
|
||||
wrapped_model, _plan = multidiffusion_sampling.clone_model_with_multidiffusion(
|
||||
model,
|
||||
latent_width=8,
|
||||
latent_height=4,
|
||||
tile_width=4,
|
||||
tile_height=4,
|
||||
overlap=0,
|
||||
tile_batch_size=2,
|
||||
differential_diffusion=True,
|
||||
)
|
||||
|
||||
assert model.clone_count == 1
|
||||
assert wrapped_model.parent is model
|
||||
assert callable(wrapped_model.model_options["denoise_mask_function"])
|
||||
assert isinstance(
|
||||
wrapped_model.wrapper,
|
||||
multidiffusion_sampling.MultiDiffusionModelWrapper,
|
||||
)
|
||||
|
||||
|
||||
def test_clone_model_rejects_non_callable_existing_wrapper() -> None:
|
||||
"""Existing wrapper metadata must be callable."""
|
||||
|
||||
|
||||
+48
-146
@@ -2,36 +2,16 @@
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Coverage tests for ComfyUI node tooltip metadata."""
|
||||
"""Coverage tests for Comfy v3 node tooltip metadata."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from collections.abc import Mapping
|
||||
from types import ModuleType
|
||||
from typing import Any, Protocol, cast
|
||||
from typing import Any, Protocol
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
from simple_syrup.nodes_v3.encode_prompt_batch_with_prompt_control import (
|
||||
EncodePromptBatchWithPromptControl,
|
||||
)
|
||||
from simple_syrup.nodes_v3.scale_factor import ScaleFactorV3
|
||||
from simple_syrup.nodes_v3.simple_load_checkpoint import SimpleLoadCheckpointV3
|
||||
from simple_syrup.nodes_v3.tile_and_tag_segs import TileAndTagSEGSV3
|
||||
from simple_syrup.nodes_v3.wd14_tagger_loader import WD14TaggerLoaderV3
|
||||
|
||||
|
||||
class _LegacyNode(Protocol):
|
||||
"""Protocol for legacy ComfyUI node declarations."""
|
||||
|
||||
DESCRIPTION: str
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]:
|
||||
"""Return ComfyUI legacy input metadata."""
|
||||
|
||||
|
||||
class _V3Node(Protocol):
|
||||
"""Protocol for Comfy v3 node schema declarations."""
|
||||
@@ -52,6 +32,9 @@ class _FakeFolderPaths(ModuleType):
|
||||
self.user_directory = "E:\\ComfyUI\\user"
|
||||
self._files = {
|
||||
"checkpoints": ["model.safetensors"],
|
||||
"diffusion_models": ["diffusion.safetensors"],
|
||||
"text_encoders": ["text_encoder.safetensors"],
|
||||
"unet": ["model.safetensors"],
|
||||
"vae": ["manual_vae.safetensors"],
|
||||
"vae_approx": [],
|
||||
}
|
||||
@@ -59,139 +42,58 @@ class _FakeFolderPaths(ModuleType):
|
||||
def get_filename_list(self, folder_name: str) -> list[str]:
|
||||
"""Return deterministic filenames for a model folder."""
|
||||
|
||||
return self._files[folder_name]
|
||||
return self._files.get(folder_name, [])
|
||||
|
||||
|
||||
def test_legacy_nodes_provide_tooltip_metadata() -> None:
|
||||
"""All exported legacy nodes expose descriptions and field-level help."""
|
||||
|
||||
for node_id, raw_node_class in NODE_CLASS_MAPPINGS.items():
|
||||
node_class = cast(_LegacyNode, raw_node_class)
|
||||
assert node_id in NODE_DISPLAY_NAME_MAPPINGS
|
||||
description = getattr(node_class, "DESCRIPTION", None)
|
||||
assert isinstance(description, str) and description.strip(), (
|
||||
f"{node_id} is missing DESCRIPTION."
|
||||
)
|
||||
|
||||
input_types = node_class.INPUT_TYPES()
|
||||
assert isinstance(input_types, Mapping), f"{node_id} INPUT_TYPES is invalid."
|
||||
for section_name in ("required", "optional", "hidden"):
|
||||
section = input_types.get(section_name, {})
|
||||
assert isinstance(section, Mapping), (
|
||||
f"{node_id} {section_name} inputs must be a mapping."
|
||||
)
|
||||
for field_name, declaration in section.items():
|
||||
if section_name == "hidden" and _legacy_hidden_sentinel(declaration):
|
||||
continue
|
||||
assert _tooltip_from_legacy_declaration(declaration), (
|
||||
f"{node_id} {section_name}.{field_name} is missing tooltip "
|
||||
"metadata."
|
||||
)
|
||||
|
||||
|
||||
def test_legacy_named_outputs_provide_tooltips() -> None:
|
||||
"""All named legacy outputs provide matching output tooltip metadata."""
|
||||
|
||||
for node_id, node_class in NODE_CLASS_MAPPINGS.items():
|
||||
return_names = getattr(node_class, "RETURN_NAMES", None)
|
||||
if return_names is None:
|
||||
continue
|
||||
|
||||
output_tooltips = getattr(node_class, "OUTPUT_TOOLTIPS", None)
|
||||
assert isinstance(output_tooltips, tuple), (
|
||||
f"{node_id} is missing OUTPUT_TOOLTIPS."
|
||||
)
|
||||
assert len(output_tooltips) == len(return_names), (
|
||||
f"{node_id} OUTPUT_TOOLTIPS must match RETURN_NAMES length."
|
||||
)
|
||||
for output_name, tooltip in zip(return_names, output_tooltips, strict=True):
|
||||
assert isinstance(tooltip, str) and tooltip.strip(), (
|
||||
f"{node_id} output.{output_name} is missing tooltip metadata."
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"schema_class",
|
||||
[
|
||||
SimpleLoadCheckpointV3,
|
||||
ScaleFactorV3,
|
||||
TileAndTagSEGSV3,
|
||||
WD14TaggerLoaderV3,
|
||||
EncodePromptBatchWithPromptControl,
|
||||
],
|
||||
)
|
||||
def test_v3_nodes_provide_tooltip_metadata(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
schema_class: type[_V3Node],
|
||||
) -> None:
|
||||
def test_v3_nodes_provide_tooltip_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""All supported v3 schemas expose descriptions and field-level help."""
|
||||
|
||||
monkeypatch.setitem(sys.modules, "folder_paths", _FakeFolderPaths())
|
||||
nodes_v3 = __import__("simple_syrup.nodes_v3", fromlist=["get_nodes"])
|
||||
monkeypatch.setattr(nodes_v3, "prompt_control_is_available", lambda: True)
|
||||
|
||||
schema = schema_class.define_schema()
|
||||
assert isinstance(schema.description, str) and schema.description.strip(), (
|
||||
f"{schema.node_id} v3 schema is missing description."
|
||||
)
|
||||
for input_item in schema.inputs:
|
||||
tooltip = getattr(input_item, "tooltip", None)
|
||||
assert isinstance(tooltip, str) and tooltip.strip(), (
|
||||
f"{schema.node_id} input.{input_item.id} is missing tooltip metadata."
|
||||
)
|
||||
for output in schema.outputs:
|
||||
tooltip = getattr(output, "tooltip", None)
|
||||
assert isinstance(tooltip, str) and tooltip.strip(), (
|
||||
f"{schema.node_id} output.{output.id} is missing tooltip metadata."
|
||||
for schema_class in nodes_v3.get_nodes():
|
||||
schema = schema_class.define_schema()
|
||||
assert isinstance(schema.description, str) and schema.description.strip(), (
|
||||
f"{schema.node_id} v3 schema is missing description."
|
||||
)
|
||||
for input_item in schema.inputs:
|
||||
tooltip = getattr(input_item, "tooltip", None)
|
||||
assert isinstance(tooltip, str) and tooltip.strip(), (
|
||||
f"{schema.node_id} input.{input_item.id} is missing tooltip metadata."
|
||||
)
|
||||
for output in schema.outputs:
|
||||
tooltip = getattr(output, "tooltip", None)
|
||||
assert isinstance(tooltip, str) and tooltip.strip(), (
|
||||
f"{schema.node_id} output.{output.id} is missing tooltip metadata."
|
||||
)
|
||||
|
||||
|
||||
def test_high_impact_tooltips_explain_direction_and_units() -> None:
|
||||
def test_high_impact_tooltips_explain_direction_and_units(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Important numeric controls explain units or practical direction."""
|
||||
|
||||
detail_node = cast(
|
||||
_LegacyNode,
|
||||
NODE_CLASS_MAPPINGS["SimpleSyrup.DetailSEGSByScaleFactor"],
|
||||
)
|
||||
detail_inputs = detail_node.INPUT_TYPES()["required"]
|
||||
denoise = _tooltip_from_legacy_declaration(detail_inputs["denoise"])
|
||||
feather = _tooltip_from_legacy_declaration(detail_inputs["feather"])
|
||||
assert "lower" in denoise.lower() and "higher" in denoise.lower()
|
||||
assert "pixels" in feather.lower()
|
||||
|
||||
tiled_node = cast(
|
||||
_LegacyNode,
|
||||
NODE_CLASS_MAPPINGS["SimpleSyrup.KSamplerTiledDiffusion"],
|
||||
)
|
||||
tiled_inputs = tiled_node.INPUT_TYPES()["required"]
|
||||
overlap = _tooltip_from_legacy_declaration(tiled_inputs["latent_tile_overlap"])
|
||||
batch_size = _tooltip_from_legacy_declaration(
|
||||
tiled_inputs["latent_tile_batch_size"]
|
||||
)
|
||||
assert "overlap" in overlap.lower() and "seams" in overlap.lower()
|
||||
assert "memory" in batch_size.lower()
|
||||
|
||||
|
||||
def _tooltip_from_legacy_declaration(declaration: object) -> str:
|
||||
"""Return a legacy ComfyUI field tooltip or an empty string."""
|
||||
|
||||
if not isinstance(declaration, tuple) or len(declaration) < 2:
|
||||
return ""
|
||||
options = declaration[1]
|
||||
if not isinstance(options, dict):
|
||||
return ""
|
||||
tooltip: Any = options.get("tooltip")
|
||||
if not isinstance(tooltip, str):
|
||||
return ""
|
||||
return tooltip.strip()
|
||||
|
||||
|
||||
def _legacy_hidden_sentinel(declaration: object) -> bool:
|
||||
"""Return whether a declaration is a Comfy legacy hidden input sentinel."""
|
||||
|
||||
return isinstance(declaration, str) and declaration in {
|
||||
"PROMPT",
|
||||
"DYNPROMPT",
|
||||
"EXTRA_PNGINFO",
|
||||
"UNIQUE_ID",
|
||||
"AUTH_TOKEN_COMFY_ORG",
|
||||
"API_KEY_COMFY_ORG",
|
||||
monkeypatch.setitem(sys.modules, "folder_paths", _FakeFolderPaths())
|
||||
nodes_v3 = __import__("simple_syrup.nodes_v3", fromlist=["get_nodes"])
|
||||
monkeypatch.setattr(nodes_v3, "prompt_control_is_available", lambda: False)
|
||||
schemas = {
|
||||
node.define_schema().node_id: node.define_schema()
|
||||
for node in nodes_v3.get_nodes()
|
||||
}
|
||||
|
||||
detail_inputs = _inputs_by_id(schemas["SimpleSyrup.DetailSEGSByScaleFactor"])
|
||||
assert "lower" in detail_inputs["denoise"].tooltip.lower()
|
||||
assert "higher" in detail_inputs["denoise"].tooltip.lower()
|
||||
assert "pixels" in detail_inputs["feather"].tooltip.lower()
|
||||
|
||||
tiled_inputs = _inputs_by_id(schemas["SimpleSyrup.KSamplerTiledDiffusion"])
|
||||
assert "overlap" in tiled_inputs["latent_tile_overlap"].tooltip.lower()
|
||||
assert "seams" in tiled_inputs["latent_tile_overlap"].tooltip.lower()
|
||||
assert "memory" in tiled_inputs["latent_tile_batch_size"].tooltip.lower()
|
||||
|
||||
|
||||
def _inputs_by_id(schema: Any) -> dict[str, Any]:
|
||||
"""Return schema inputs keyed by id."""
|
||||
|
||||
return {input_item.id: input_item for input_item in schema.inputs}
|
||||
|
||||
@@ -24,6 +24,7 @@ EXPECTED_RUNTIME_REQUIREMENTS = (
|
||||
"addict",
|
||||
"yapf",
|
||||
"huggingface-hub",
|
||||
"keyring",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for Prompt-Control prompt preparation helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.domain.prompt_control_prompt import (
|
||||
apply_encode_style,
|
||||
extract_lora_tags,
|
||||
extract_prompt_text,
|
||||
prepare_prompt_side,
|
||||
)
|
||||
|
||||
|
||||
def test_extract_prompt_text_matches_sugarcubes_regex_behavior() -> None:
|
||||
"""Visible text outside angle tags is retained and joined with newlines."""
|
||||
|
||||
text = "portrait <lora:face:1.0> cinematic <foo>"
|
||||
|
||||
assert extract_prompt_text(text) == "portrait \n cinematic "
|
||||
|
||||
|
||||
def test_extract_prompt_text_keeps_text_after_adjacent_tag() -> None:
|
||||
"""Text after a closing angle tag is captured as a later prompt segment."""
|
||||
|
||||
assert extract_prompt_text("<lora:a:1.0>face") == "face"
|
||||
|
||||
|
||||
def test_extract_lora_tags_returns_angle_tags_joined_with_newlines() -> None:
|
||||
"""All Prompt-Control angle tags are retained for LoRA scheduling."""
|
||||
|
||||
text = "portrait <lora:face:1.0> cinematic <lora:light:0.5>"
|
||||
|
||||
assert extract_lora_tags(text) == "<lora:face:1.0>\n<lora:light:0.5>"
|
||||
|
||||
|
||||
def test_prepare_prompt_side_splits_ordered_chunks_and_aggregates_loras() -> None:
|
||||
"""Separator-delimited chunks preserve order and collect all tags."""
|
||||
|
||||
side = prepare_prompt_side(
|
||||
"face <lora:a:1.0> [SEP] hair <lora:b:0.5>",
|
||||
"[SEP]",
|
||||
)
|
||||
|
||||
assert [chunk.text for chunk in side.chunks] == ["face ", "hair "]
|
||||
assert [chunk.lora_tags for chunk in side.chunks] == [
|
||||
"<lora:a:1.0>",
|
||||
"<lora:b:0.5>",
|
||||
]
|
||||
assert side.lora_tags == "<lora:a:1.0>\n<lora:b:0.5>"
|
||||
|
||||
|
||||
def test_prepare_prompt_side_preserves_empty_chunks() -> None:
|
||||
"""Batch splitting keeps empty entries so batch positions stay explicit."""
|
||||
|
||||
side = prepare_prompt_side("face [SEP] ", "[SEP]")
|
||||
|
||||
assert [chunk.text for chunk in side.chunks] == ["face", ""]
|
||||
assert [chunk.lora_tags for chunk in side.chunks] == ["", ""]
|
||||
assert side.lora_tags == ""
|
||||
|
||||
|
||||
def test_prepare_prompt_side_rejects_empty_separator() -> None:
|
||||
"""Empty separators are rejected by the shared batch splitter."""
|
||||
|
||||
with pytest.raises(ValueError, match="separator must not be empty"):
|
||||
prepare_prompt_side("face", "")
|
||||
|
||||
|
||||
def test_apply_encode_style_prepends_style_without_extra_formatting() -> None:
|
||||
"""Encode style text is used exactly as produced by style nodes."""
|
||||
|
||||
assert apply_encode_style("STYLE(A1111) ", "face") == "STYLE(A1111) face"
|
||||
|
||||
|
||||
def test_apply_encode_style_keeps_prompt_when_style_is_blank() -> None:
|
||||
"""Blank encode style leaves cleaned prompt text unchanged."""
|
||||
|
||||
assert apply_encode_style("", "face") == "face"
|
||||
@@ -0,0 +1,224 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for Prompt-Control schedule and encode lazy graph expansion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from importlib import import_module
|
||||
from types import ModuleType
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from simple_syrup.runtime.prompt_control_schedule_encode_graph import (
|
||||
PROMPT_CONTROL_MISSING_MESSAGE,
|
||||
PromptControlScheduleEncodeGraphBuilder,
|
||||
)
|
||||
|
||||
|
||||
def test_schedule_encode_graph_builds_single_conditioning_outputs(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Single prompts return direct Prompt-Control conditioning links."""
|
||||
|
||||
calls = _install_fake_prompt_control(monkeypatch)
|
||||
|
||||
output = PromptControlScheduleEncodeGraphBuilder().build(
|
||||
model=["model", 0],
|
||||
clip=["clip", 0],
|
||||
positive_prompt="face <lora:positive:1.0>",
|
||||
negative_prompt="blur <lora:negative:0.5>",
|
||||
encode_style="STYLE(A1111) ",
|
||||
)
|
||||
|
||||
assert output.args[0] == ["lora_negative", 0]
|
||||
assert output.args[1] == ["encode_0", 0]
|
||||
assert output.args[2] == ["encode_1", 0]
|
||||
assert output.expand is not None
|
||||
assert not any(
|
||||
node["class_type"].startswith("SimpleSyrup.ConditioningBatch")
|
||||
for node in output.expand.values()
|
||||
)
|
||||
assert calls["lora"] == [
|
||||
{
|
||||
"model": ["model", 0],
|
||||
"clip": ["clip", 0],
|
||||
"text": "<lora:positive:1.0>",
|
||||
},
|
||||
{
|
||||
"model": ["lora_positive", 0],
|
||||
"clip": ["lora_positive", 1],
|
||||
"text": "<lora:negative:0.5>",
|
||||
},
|
||||
]
|
||||
assert calls["encode"] == [
|
||||
{"clip": ["lora_negative", 1], "text": "STYLE(A1111) face "},
|
||||
{"clip": ["lora_negative", 1], "text": "STYLE(A1111) blur "},
|
||||
]
|
||||
|
||||
|
||||
def test_schedule_encode_graph_packs_only_multichunk_sides(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""A side with multiple chunks becomes a conditioning batch."""
|
||||
|
||||
_install_fake_prompt_control(monkeypatch)
|
||||
|
||||
output = PromptControlScheduleEncodeGraphBuilder().build(
|
||||
model=["model", 0],
|
||||
clip=["clip", 0],
|
||||
positive_prompt="face [SEP] hair",
|
||||
negative_prompt="blur",
|
||||
)
|
||||
|
||||
assert output.args[1] != ["encode_0", 0]
|
||||
assert output.args[2] == ["encode_2", 0]
|
||||
assert output.expand is not None
|
||||
pack_nodes = [
|
||||
node
|
||||
for node in output.expand.values()
|
||||
if node["class_type"].startswith("SimpleSyrup.ConditioningBatch")
|
||||
]
|
||||
assert [node["class_type"] for node in pack_nodes] == [
|
||||
"SimpleSyrup.ConditioningBatchStart",
|
||||
"SimpleSyrup.ConditioningBatchAppend",
|
||||
]
|
||||
|
||||
|
||||
def test_schedule_encode_graph_collects_loras_from_all_chunks(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""LoRA scheduling sees all tags from every separator-delimited chunk."""
|
||||
|
||||
calls = _install_fake_prompt_control(monkeypatch)
|
||||
|
||||
PromptControlScheduleEncodeGraphBuilder().build(
|
||||
model=["model", 0],
|
||||
clip=["clip", 0],
|
||||
positive_prompt="face <lora:a:1> [SEP] hair <lora:b:1>",
|
||||
negative_prompt="blur <lora:c:1> [SEP] noise <lora:d:1>",
|
||||
)
|
||||
|
||||
assert calls["lora"][0]["text"] == "<lora:a:1>\n<lora:b:1>"
|
||||
assert calls["lora"][1]["text"] == "<lora:c:1>\n<lora:d:1>"
|
||||
|
||||
|
||||
def test_schedule_encode_graph_reports_duplicate_expand_ids(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Duplicate generated node ids are rejected with context."""
|
||||
|
||||
_install_fake_prompt_control(monkeypatch, duplicate_lora_ids=True)
|
||||
|
||||
with pytest.raises(ValueError, match="negative LoRA scheduling"):
|
||||
PromptControlScheduleEncodeGraphBuilder().build(
|
||||
model=["model", 0],
|
||||
clip=["clip", 0],
|
||||
positive_prompt="face",
|
||||
negative_prompt="blur",
|
||||
)
|
||||
|
||||
|
||||
def test_schedule_encode_graph_reports_missing_prompt_control(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Missing Prompt-Control dependency raises an actionable error."""
|
||||
|
||||
def fake_import_module(name: str) -> Any:
|
||||
if name == "prompt_control.nodes_lazy":
|
||||
raise ModuleNotFoundError(name)
|
||||
return import_module(name)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"simple_syrup.runtime.prompt_control_schedule_encode_graph.import_module",
|
||||
fake_import_module,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="requires comfyui-prompt-control"):
|
||||
PromptControlScheduleEncodeGraphBuilder().build(
|
||||
model=["model", 0],
|
||||
clip=["clip", 0],
|
||||
positive_prompt="face",
|
||||
negative_prompt="blur",
|
||||
)
|
||||
assert PROMPT_CONTROL_MISSING_MESSAGE.startswith("Schedule & Encode Prompts")
|
||||
|
||||
|
||||
def _install_fake_prompt_control(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
duplicate_lora_ids: bool = False,
|
||||
) -> dict[str, list[dict[str, Any]]]:
|
||||
"""Install graph-expanding Prompt-Control stand-ins for builder tests."""
|
||||
|
||||
prompt_control = ModuleType("prompt_control")
|
||||
nodes_lazy = ModuleType("prompt_control.nodes_lazy")
|
||||
io = import_module("comfy_api.latest").io
|
||||
calls: dict[str, list[dict[str, Any]]] = {"lora": [], "encode": []}
|
||||
|
||||
class FakePCLazyLoraLoaderAdvanced:
|
||||
"""Prompt-Control LoRA scheduler test double."""
|
||||
|
||||
@staticmethod
|
||||
def execute(
|
||||
model: Any,
|
||||
clip: Any,
|
||||
text: str,
|
||||
apply_hooks: bool,
|
||||
tags: str,
|
||||
start: float,
|
||||
end: float,
|
||||
num_steps: int,
|
||||
) -> Any:
|
||||
"""Return deterministic model and clip links."""
|
||||
|
||||
assert apply_hooks is True
|
||||
assert tags == ""
|
||||
assert start == 0.0
|
||||
assert end == 1.0
|
||||
assert num_steps == 0
|
||||
side = "positive" if not calls["lora"] else "negative"
|
||||
calls["lora"].append({"model": model, "clip": clip, "text": text})
|
||||
node_id = "duplicate_lora" if duplicate_lora_ids else f"lora_{side}"
|
||||
return io.NodeOutput(
|
||||
[f"lora_{side}", 0],
|
||||
[f"lora_{side}", 1],
|
||||
None,
|
||||
expand={node_id: {"class_type": "PromptControl.FakeLora"}},
|
||||
)
|
||||
|
||||
class FakePCLazyTextEncodeAdvanced:
|
||||
"""Prompt-Control text encoder test double."""
|
||||
|
||||
@staticmethod
|
||||
def execute(
|
||||
clip: Any,
|
||||
text: str,
|
||||
tags: str,
|
||||
start: float,
|
||||
end: float,
|
||||
num_steps: int,
|
||||
) -> Any:
|
||||
"""Return deterministic conditioning links."""
|
||||
|
||||
assert tags == ""
|
||||
assert start == 0.0
|
||||
assert end == 1.0
|
||||
assert num_steps == 0
|
||||
index = len(calls["encode"])
|
||||
calls["encode"].append({"clip": clip, "text": text})
|
||||
node_id = f"encode_{index}"
|
||||
return io.NodeOutput(
|
||||
[node_id, 0],
|
||||
expand={node_id: {"class_type": "PromptControl.FakeTextEncode"}},
|
||||
)
|
||||
|
||||
cast(Any, nodes_lazy).PCLazyLoraLoaderAdvanced = FakePCLazyLoraLoaderAdvanced
|
||||
cast(Any, nodes_lazy).PCLazyTextEncodeAdvanced = FakePCLazyTextEncodeAdvanced
|
||||
cast(Any, prompt_control).nodes_lazy = nodes_lazy
|
||||
monkeypatch.setitem(sys.modules, "prompt_control", prompt_control)
|
||||
monkeypatch.setitem(sys.modules, "prompt_control.nodes_lazy", nodes_lazy)
|
||||
return calls
|
||||
@@ -29,7 +29,7 @@ def test_prompt_encode_style_node_contract_constants() -> None:
|
||||
"""Style-only node constants match the public ComfyUI contract."""
|
||||
|
||||
assert PromptEncodeStyle.RETURN_TYPES == ("STRING",)
|
||||
assert PromptEncodeStyle.RETURN_NAMES == ("style_tag",)
|
||||
assert PromptEncodeStyle.RETURN_NAMES == ("encode_style",)
|
||||
assert PromptEncodeStyle.FUNCTION == "build"
|
||||
assert PromptEncodeStyle.CATEGORY == "SimpleSyrup/Prompt"
|
||||
|
||||
@@ -60,18 +60,18 @@ def test_prompt_encode_style_node_builds_style_tag(
|
||||
) -> None:
|
||||
"""Style-only node formats Prompt Control STYLE tags from combo values."""
|
||||
|
||||
(style_tag,) = PromptEncodeStyle().build(encode_style)
|
||||
(encode_style_text,) = PromptEncodeStyle().build(encode_style)
|
||||
|
||||
assert style_tag == expected
|
||||
assert style_tag.endswith(" ")
|
||||
assert not style_tag.endswith(" ")
|
||||
assert encode_style_text == expected
|
||||
assert encode_style_text.endswith(" ")
|
||||
assert not encode_style_text.endswith(" ")
|
||||
|
||||
|
||||
def test_prompt_encode_style_and_normalization_node_contract_constants() -> None:
|
||||
"""Normalization node constants match the public ComfyUI contract."""
|
||||
|
||||
assert PromptEncodeStyleAndNormalization.RETURN_TYPES == ("STRING",)
|
||||
assert PromptEncodeStyleAndNormalization.RETURN_NAMES == ("style_tag",)
|
||||
assert PromptEncodeStyleAndNormalization.RETURN_NAMES == ("encode_style",)
|
||||
assert PromptEncodeStyleAndNormalization.FUNCTION == "build"
|
||||
assert PromptEncodeStyleAndNormalization.CATEGORY == "SimpleSyrup/Prompt"
|
||||
|
||||
@@ -110,10 +110,10 @@ def test_prompt_encode_style_and_normalization_node_builds_style_tag(
|
||||
) -> None:
|
||||
"""Normalization node formats Prompt Control STYLE tags from combo values."""
|
||||
|
||||
(style_tag,) = PromptEncodeStyleAndNormalization().build(
|
||||
(encode_style_text,) = PromptEncodeStyleAndNormalization().build(
|
||||
encode_style, normalization
|
||||
)
|
||||
|
||||
assert style_tag == expected
|
||||
assert style_tag.endswith(" ")
|
||||
assert not style_tag.endswith(" ")
|
||||
assert encode_style_text == expected
|
||||
assert encode_style_text.endswith(" ")
|
||||
assert not encode_style_text.endswith(" ")
|
||||
|
||||
@@ -94,6 +94,7 @@ class FakeModel:
|
||||
def __init__(
|
||||
self,
|
||||
model_options: dict[str, Any] | None = None,
|
||||
parent: FakeModel | None = None,
|
||||
) -> None:
|
||||
"""Create a fake model patcher."""
|
||||
|
||||
@@ -101,11 +102,14 @@ class FakeModel:
|
||||
self.model_options = {} if model_options is None else model_options
|
||||
self.calc_wrapper: Any = None
|
||||
self.model_sampling = object()
|
||||
self.parent = parent
|
||||
self.clone_count = 0
|
||||
|
||||
def clone(self) -> FakeModel:
|
||||
"""Return a cloned model with copied options."""
|
||||
|
||||
return FakeModel(self.model_options.copy())
|
||||
self.clone_count += 1
|
||||
return FakeModel(self.model_options.copy(), parent=self)
|
||||
|
||||
def set_model_sampler_calc_cond_batch_function(self, wrapper: object) -> None:
|
||||
"""Capture the installed calc-cond-batch wrapper."""
|
||||
@@ -113,6 +117,11 @@ class FakeModel:
|
||||
self.calc_wrapper = wrapper
|
||||
self.model_options["sampler_calc_cond_batch_function"] = wrapper
|
||||
|
||||
def set_model_denoise_mask_function(self, denoise_mask_function: object) -> None:
|
||||
"""Capture the installed denoise-mask function."""
|
||||
|
||||
self.model_options["denoise_mask_function"] = denoise_mask_function
|
||||
|
||||
def get_model_object(self, name: str) -> object:
|
||||
"""Return the requested fake model object."""
|
||||
|
||||
@@ -153,6 +162,31 @@ def test_clone_model_installs_regional_calc_cond_batch_wrapper() -> None:
|
||||
assert summary.region_count == 1
|
||||
|
||||
|
||||
def test_clone_model_composes_differential_on_same_clone() -> None:
|
||||
"""Differential diffusion is installed without cloning a temporary parent."""
|
||||
|
||||
model = FakeModel()
|
||||
|
||||
wrapped_model, _summary = (
|
||||
regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion(
|
||||
model,
|
||||
latent_width=8,
|
||||
latent_height=4,
|
||||
latent_ndim=4,
|
||||
regions=(_region(0, 0, 4, 4, "positive", latent_width=8),),
|
||||
differential_diffusion=True,
|
||||
)
|
||||
)
|
||||
|
||||
assert model.clone_count == 1
|
||||
assert wrapped_model.parent is model
|
||||
assert callable(wrapped_model.model_options["denoise_mask_function"])
|
||||
assert isinstance(
|
||||
wrapped_model.calc_wrapper,
|
||||
regional_multidiffusion_sampling.RegionalMultiDiffusionCalcCondBatch,
|
||||
)
|
||||
|
||||
|
||||
def test_clone_model_rejects_non_callable_existing_calc_wrapper() -> None:
|
||||
"""Existing calc-cond-batch metadata must be callable."""
|
||||
|
||||
|
||||
+105
-429
@@ -2,7 +2,7 @@
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for SimpleSyrup ComfyUI node registration."""
|
||||
"""Tests for SimpleSyrup Comfy v3-only node registration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -12,19 +12,81 @@ import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
import pytest
|
||||
|
||||
BASE_NODE_IDS = [
|
||||
"SimpleSyrup.BatchRegionConditioning",
|
||||
"SimpleSyrup.BatchSEGS",
|
||||
"SimpleSyrup.ConditioningBatchAppend",
|
||||
"SimpleSyrup.ConditioningBatchStart",
|
||||
"SimpleSyrup.DetailSEGSAsRegions",
|
||||
"SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion",
|
||||
"SimpleSyrup.DetailSEGSByScaleFactor",
|
||||
"SimpleSyrup.DetectSEGSWithUltralytics",
|
||||
"SimpleSyrup.EncodePromptBatch",
|
||||
"SimpleSyrup.ExternalLLMPrompt",
|
||||
"SimpleSyrup.GroundedSAMModelInfo",
|
||||
"SimpleSyrup.GroundingDINOModelLoader",
|
||||
"SimpleSyrup.KSamplerExtras",
|
||||
"SimpleSyrup.KSamplerTiledDiffusion",
|
||||
"SimpleSyrup.LatentDiagnostics",
|
||||
"SimpleSyrup.LayerStyleSAMModelsAdapter",
|
||||
"SimpleSyrup.LoadUltralyticsModel",
|
||||
"SimpleSyrup.PromptEncodeStyleAndNormalization",
|
||||
"SimpleSyrup.PromptEncodeStyle",
|
||||
"SimpleSyrup.PromptSEGSWithSAM",
|
||||
"SimpleSyrup.ResizeImageToTarget",
|
||||
"SimpleSyrup.SAMModelLoader",
|
||||
"SimpleSyrup.ScaleFactor",
|
||||
"SimpleSyrup.Seed",
|
||||
"SimpleSyrup.SimpleLoadAnima",
|
||||
"SimpleSyrup.SimpleLoadCheckpoint",
|
||||
"SimpleSyrup.SimpleVAEEncode",
|
||||
"SimpleSyrup.TagSEGSWithExternalLLM",
|
||||
"SimpleSyrup.TagSEGSWithWD14",
|
||||
"SimpleSyrup.TileAndTagSEGS",
|
||||
"SimpleSyrup.UpscaleLatentFromImage",
|
||||
"SimpleSyrup.VAEDecodeOptions",
|
||||
"SimpleSyrup.VAEEncodeOptions",
|
||||
"SimpleSyrup.ViTMatteModelLoader",
|
||||
"SimpleSyrup.WD14TaggerLoader",
|
||||
]
|
||||
|
||||
def test_package_exports_node_mappings() -> None:
|
||||
"""Root package import exposes ComfyUI mapping dictionaries."""
|
||||
PROMPT_CONTROL_NODE_IDS = [
|
||||
"SimpleSyrup.EncodePromptBatchWithPromptControl",
|
||||
"SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl",
|
||||
]
|
||||
|
||||
|
||||
class _V3Node(Protocol):
|
||||
"""Protocol for Comfy v3 node schema declarations."""
|
||||
|
||||
@classmethod
|
||||
def define_schema(cls) -> Any:
|
||||
"""Return a Comfy v3 schema object."""
|
||||
|
||||
|
||||
def test_package_exports_v3_entrypoint_only() -> None:
|
||||
"""Root package exposes Comfy v3 registration and no legacy mappings."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
|
||||
assert hasattr(package, "NODE_CLASS_MAPPINGS")
|
||||
assert hasattr(package, "NODE_DISPLAY_NAME_MAPPINGS")
|
||||
assert hasattr(package, "comfy_entrypoint")
|
||||
assert package.WEB_DIRECTORY == "./web/dist"
|
||||
assert not hasattr(package, "NODE_CLASS_MAPPINGS")
|
||||
assert not hasattr(package, "NODE_DISPLAY_NAME_MAPPINGS")
|
||||
assert package.__all__ == ["WEB_DIRECTORY", "comfy_entrypoint"]
|
||||
|
||||
|
||||
def test_nodes_package_is_not_a_legacy_registry() -> None:
|
||||
"""The implementation package no longer owns ComfyUI registration."""
|
||||
|
||||
nodes_package = importlib.import_module("SimpleSyrup.simple_syrup.nodes")
|
||||
|
||||
assert not hasattr(nodes_package, "NODE_CLASS_MAPPINGS")
|
||||
assert not hasattr(nodes_package, "NODE_DISPLAY_NAME_MAPPINGS")
|
||||
|
||||
|
||||
def test_package_imports_from_custom_nodes_parent_path() -> None:
|
||||
@@ -39,7 +101,9 @@ def test_package_imports_from_custom_nodes_parent_path() -> None:
|
||||
"if pathlib.Path(p or '.').resolve() != project]; "
|
||||
f"sys.path.insert(0, {str(custom_nodes_root)!r}); "
|
||||
"package = importlib.import_module('SimpleSyrup'); "
|
||||
"assert 'SimpleSyrup.PromptSEGSWithSAM' in package.NODE_CLASS_MAPPINGS; "
|
||||
"assert hasattr(package, 'comfy_entrypoint'); "
|
||||
"assert not hasattr(package, 'NODE_CLASS_MAPPINGS'); "
|
||||
"assert not hasattr(package, 'NODE_DISPLAY_NAME_MAPPINGS'); "
|
||||
"assert 'server' not in sys.modules"
|
||||
)
|
||||
|
||||
@@ -83,399 +147,6 @@ def test_comfy_import_exposes_stable_internal_package_alias() -> None:
|
||||
assert result.returncode == 0, result.stderr
|
||||
|
||||
|
||||
def test_resize_node_is_registered() -> None:
|
||||
"""Resize node id maps to the expected node class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.ResizeImageToTarget"]
|
||||
|
||||
assert registered.__name__ == "ResizeImageToTarget"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.ResizeImageToTarget"]
|
||||
== "Resize Image to Target"
|
||||
)
|
||||
|
||||
|
||||
def test_ksampler_extras_node_is_registered() -> None:
|
||||
"""KSampler Extras node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.KSamplerExtras"]
|
||||
|
||||
assert registered.__name__ == "KSamplerExtras"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.KSamplerExtras"]
|
||||
== "KSampler (Extras)"
|
||||
)
|
||||
|
||||
|
||||
def test_ksampler_tiled_diffusion_node_is_registered() -> None:
|
||||
"""KSampler tiled diffusion node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.KSamplerTiledDiffusion"]
|
||||
|
||||
assert registered.__name__ == "KSamplerTiledDiffusion"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.KSamplerTiledDiffusion"]
|
||||
== "KSampler (Tiled Diffusion)"
|
||||
)
|
||||
assert (
|
||||
"KSamplerTiledDiffusion"
|
||||
in importlib.import_module("SimpleSyrup.simple_syrup.nodes").__all__
|
||||
)
|
||||
assert "SimpleSyrup.KSamplerMixtureOfDiffusers" not in package.NODE_CLASS_MAPPINGS
|
||||
assert "SimpleSyrup.KSamplerMultiDiffusion" not in package.NODE_CLASS_MAPPINGS
|
||||
assert (
|
||||
"SimpleSyrup.KSamplerMixtureOfDiffusers"
|
||||
not in package.NODE_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
assert (
|
||||
"SimpleSyrup.KSamplerMultiDiffusion" not in package.NODE_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
|
||||
|
||||
def test_latent_diagnostics_node_is_registered() -> None:
|
||||
"""Latent Diagnostics node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.LatentDiagnostics"]
|
||||
|
||||
assert registered.__name__ == "LatentDiagnostics"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.LatentDiagnostics"]
|
||||
== "Latent Diagnostics"
|
||||
)
|
||||
|
||||
|
||||
def test_prompt_encode_style_node_is_registered() -> None:
|
||||
"""Prompt Encode Style node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.PromptEncodeStyle"]
|
||||
|
||||
assert registered.__name__ == "PromptEncodeStyle"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.PromptEncodeStyle"]
|
||||
== "Prompt Encode Style"
|
||||
)
|
||||
|
||||
|
||||
def test_prompt_encode_style_and_normalization_node_is_registered() -> None:
|
||||
"""Prompt Encode Style & Normalization node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS[
|
||||
"SimpleSyrup.PromptEncodeStyleAndNormalization"
|
||||
]
|
||||
|
||||
assert registered.__name__ == "PromptEncodeStyleAndNormalization"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS[
|
||||
"SimpleSyrup.PromptEncodeStyleAndNormalization"
|
||||
]
|
||||
== "Prompt Encode Style & Normalization"
|
||||
)
|
||||
|
||||
|
||||
def test_prompt_control_encode_style_clean_break_id_is_removed() -> None:
|
||||
"""Old Prompt Control Encode Style node id is not registered."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
|
||||
assert "SimpleSyrup.PromptControlEncodeStyle" not in package.NODE_CLASS_MAPPINGS
|
||||
assert (
|
||||
"SimpleSyrup.PromptControlEncodeStyle" not in package.NODE_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
|
||||
|
||||
def test_prompt_segs_with_sam_node_is_registered() -> None:
|
||||
"""Prompt SEGS w/ SAM node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.PromptSEGSWithSAM"]
|
||||
|
||||
assert registered.__name__ == "PromptSEGSWithSAM"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.PromptSEGSWithSAM"]
|
||||
== "Prompt SEGS w/ SAM"
|
||||
)
|
||||
assert "SimpleSyrup.PromptSAMMask" not in package.NODE_CLASS_MAPPINGS
|
||||
assert "SimpleSyrup.PromptSAMMask" not in package.NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
|
||||
def test_sam_model_loader_node_is_registered() -> None:
|
||||
"""SAM loader node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.SAMModelLoader"]
|
||||
|
||||
assert registered.__name__ == "SAMModelLoader"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.SAMModelLoader"]
|
||||
== "SAM Model Loader"
|
||||
)
|
||||
|
||||
|
||||
def test_scale_factor_node_is_registered() -> None:
|
||||
"""Scale Factor node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.ScaleFactor"]
|
||||
|
||||
assert registered.__name__ == "ScaleFactor"
|
||||
assert package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.ScaleFactor"] == (
|
||||
"Scale Factor"
|
||||
)
|
||||
assert (
|
||||
"ScaleFactor"
|
||||
in importlib.import_module("SimpleSyrup.simple_syrup.nodes").__all__
|
||||
)
|
||||
|
||||
|
||||
def test_seed_node_is_registered() -> None:
|
||||
"""Seed node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.Seed"]
|
||||
|
||||
assert registered.__name__ == "Seed"
|
||||
assert package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.Seed"] == "Seed"
|
||||
|
||||
|
||||
def test_grounding_dino_model_loader_node_is_registered() -> None:
|
||||
"""GroundingDINO loader node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.GroundingDINOModelLoader"]
|
||||
|
||||
assert registered.__name__ == "GroundingDINOModelLoader"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.GroundingDINOModelLoader"]
|
||||
== "GroundingDINO Model Loader"
|
||||
)
|
||||
|
||||
|
||||
def test_vitmatte_model_loader_node_is_registered() -> None:
|
||||
"""ViTMatte loader node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.ViTMatteModelLoader"]
|
||||
|
||||
assert registered.__name__ == "ViTMatteModelLoader"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.ViTMatteModelLoader"]
|
||||
== "ViTMatte Model Loader"
|
||||
)
|
||||
|
||||
|
||||
def test_wd14_tagger_loader_node_is_registered() -> None:
|
||||
"""WD14 tagger loader node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.WD14TaggerLoader"]
|
||||
|
||||
assert registered.__name__ == "WD14TaggerLoader"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.WD14TaggerLoader"]
|
||||
== "Load WD14 Tagger"
|
||||
)
|
||||
|
||||
|
||||
def test_load_ultralytics_model_node_is_registered() -> None:
|
||||
"""Load Ultralytics Model node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.LoadUltralyticsModel"]
|
||||
|
||||
assert registered.__name__ == "LoadUltralyticsModel"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.LoadUltralyticsModel"]
|
||||
== "Load Ultralytics Model"
|
||||
)
|
||||
|
||||
|
||||
def test_detect_segs_with_ultralytics_node_is_registered() -> None:
|
||||
"""Detect SEGS w/ Ultralytics node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.DetectSEGSWithUltralytics"]
|
||||
|
||||
assert registered.__name__ == "DetectSEGSWithUltralytics"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.DetectSEGSWithUltralytics"]
|
||||
== "Detect SEGS w/ Ultralytics"
|
||||
)
|
||||
|
||||
|
||||
def test_detail_segs_by_scale_factor_node_is_registered() -> None:
|
||||
"""Detail SEGS by Scale Factor node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.DetailSEGSByScaleFactor"]
|
||||
|
||||
assert registered.__name__ == "DetailSEGSByScaleFactor"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.DetailSEGSByScaleFactor"]
|
||||
== "Detail SEGS by Scale Factor"
|
||||
)
|
||||
|
||||
|
||||
def test_tiled_detail_segs_by_scale_factor_node_is_registered() -> None:
|
||||
"""Tiled Detail SEGS by Scale Factor node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS[
|
||||
"SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion"
|
||||
]
|
||||
|
||||
assert registered.__name__ == "DetailSEGSByScaleFactorTiledDiffusion"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS[
|
||||
"SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion"
|
||||
]
|
||||
== "Detail SEGS by Scale Factor w/ Tiled Diffusion"
|
||||
)
|
||||
assert (
|
||||
"DetailSEGSByScaleFactorTiledDiffusion"
|
||||
in importlib.import_module("SimpleSyrup.simple_syrup.nodes").__all__
|
||||
)
|
||||
|
||||
|
||||
def test_detail_segs_as_regions_node_is_registered() -> None:
|
||||
"""Detail SEGS as Regions node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.DetailSEGSAsRegions"]
|
||||
|
||||
assert registered.__name__ == "DetailSEGSAsRegions"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.DetailSEGSAsRegions"]
|
||||
== "Detail SEGS as Regions"
|
||||
)
|
||||
|
||||
|
||||
def test_tile_and_tag_segs_node_is_registered() -> None:
|
||||
"""Tile & Tag SEGS node maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.TileAndTagSEGS"]
|
||||
|
||||
assert registered.__name__ == "TileAndTagSEGS"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.TileAndTagSEGS"]
|
||||
== "Tile & Tag SEGS"
|
||||
)
|
||||
|
||||
|
||||
def test_conditioning_batch_nodes_are_registered() -> None:
|
||||
"""Conditioning batch nodes map to their classes and display names."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
|
||||
assert (
|
||||
package.NODE_CLASS_MAPPINGS["SimpleSyrup.ConditioningBatchStart"].__name__
|
||||
== "ConditioningBatchStart"
|
||||
)
|
||||
assert (
|
||||
package.NODE_CLASS_MAPPINGS["SimpleSyrup.ConditioningBatchAppend"].__name__
|
||||
== "ConditioningBatchAppend"
|
||||
)
|
||||
assert (
|
||||
package.NODE_CLASS_MAPPINGS["SimpleSyrup.EncodePromptBatch"].__name__
|
||||
== "EncodePromptBatch"
|
||||
)
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.ConditioningBatchStart"]
|
||||
== "Conditioning Batch Start"
|
||||
)
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.ConditioningBatchAppend"]
|
||||
== "Conditioning Batch Append"
|
||||
)
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.EncodePromptBatch"]
|
||||
== "Encode Prompt Batch"
|
||||
)
|
||||
|
||||
|
||||
def test_layerstyle_adapter_node_is_registered() -> None:
|
||||
"""LayerStyle adapter node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.LayerStyleSAMModelsAdapter"]
|
||||
|
||||
assert registered.__name__ == "LayerStyleSAMModelsAdapter"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.LayerStyleSAMModelsAdapter"]
|
||||
== "LayerStyle SAM Models Adapter"
|
||||
)
|
||||
|
||||
|
||||
def test_grounded_sam_model_info_node_is_registered() -> None:
|
||||
"""Grounded SAM model info node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.GroundedSAMModelInfo"]
|
||||
|
||||
assert registered.__name__ == "GroundedSAMModelInfo"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.GroundedSAMModelInfo"]
|
||||
== "Grounded SAM Model Info"
|
||||
)
|
||||
|
||||
|
||||
def test_simple_load_anima_node_is_registered() -> None:
|
||||
"""Simple Load Anima node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.SimpleLoadAnima"]
|
||||
|
||||
assert registered.__name__ == "SimpleLoadAnima"
|
||||
assert registered.RETURN_TYPES == ("MODEL", "CLIP", "VAE")
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.SimpleLoadAnima"]
|
||||
== "Simple Load Anima"
|
||||
)
|
||||
|
||||
|
||||
def test_simple_load_checkpoint_node_is_registered() -> None:
|
||||
"""Simple Load Checkpoint node id maps to its class and display name."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
registered = package.NODE_CLASS_MAPPINGS["SimpleSyrup.SimpleLoadCheckpoint"]
|
||||
|
||||
assert registered.__name__ == "SimpleLoadCheckpoint"
|
||||
assert registered.RETURN_TYPES == ("MODEL", "CLIP", "VAE")
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.SimpleLoadCheckpoint"]
|
||||
== "Simple Load Checkpoint"
|
||||
)
|
||||
|
||||
|
||||
def test_provenance_latent_nodes_are_registered() -> None:
|
||||
"""Provenance-aware latent nodes map to their classes and display names."""
|
||||
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
nodes_package = importlib.import_module("SimpleSyrup.simple_syrup.nodes")
|
||||
|
||||
simple_vae = package.NODE_CLASS_MAPPINGS["SimpleSyrup.SimpleVAEEncode"]
|
||||
upscale = package.NODE_CLASS_MAPPINGS["SimpleSyrup.UpscaleLatentFromImage"]
|
||||
|
||||
assert simple_vae.__name__ == "SimpleVAEEncode"
|
||||
assert upscale.__name__ == "UpscaleLatentFromImage"
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.SimpleVAEEncode"]
|
||||
== "Simple VAE Encode"
|
||||
)
|
||||
assert (
|
||||
package.NODE_DISPLAY_NAME_MAPPINGS["SimpleSyrup.UpscaleLatentFromImage"]
|
||||
== "Upscale Latent From Image"
|
||||
)
|
||||
assert "SimpleVAEEncode" in nodes_package.__all__
|
||||
assert "UpscaleLatentFromImage" in nodes_package.__all__
|
||||
|
||||
|
||||
def test_registration_import_does_not_require_torchlanc() -> None:
|
||||
"""Importing registration does not eagerly import TorchLanc."""
|
||||
|
||||
@@ -486,33 +157,10 @@ def test_registration_import_does_not_require_torchlanc() -> None:
|
||||
assert imported_module is None
|
||||
|
||||
|
||||
def test_v3_entrypoint_registers_tile_and_prompt_control_batch_nodes(
|
||||
def test_v3_entrypoint_exports_all_base_nodes_without_prompt_control(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Comfy v3 entrypoint exposes native v3 nodes without Prompt Control imports."""
|
||||
|
||||
sys.modules.pop("prompt_control.nodes_lazy", None)
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
nodes_v3 = importlib.import_module("SimpleSyrup.simple_syrup.nodes_v3")
|
||||
monkeypatch.setattr(nodes_v3, "prompt_control_is_available", lambda: True)
|
||||
|
||||
extension = asyncio.run(package.comfy_entrypoint())
|
||||
nodes = asyncio.run(extension.get_node_list())
|
||||
|
||||
assert [node.__name__ for node in nodes] == [
|
||||
"WD14TaggerLoaderV3",
|
||||
"TileAndTagSEGSV3",
|
||||
"SimpleLoadCheckpointV3",
|
||||
"ScaleFactorV3",
|
||||
"EncodePromptBatchWithPromptControl",
|
||||
]
|
||||
assert "prompt_control.nodes_lazy" not in sys.modules
|
||||
|
||||
|
||||
def test_v3_entrypoint_keeps_tile_node_when_prompt_control_unavailable(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Comfy v3 entrypoint omits only Prompt Control nodes when unavailable."""
|
||||
"""Comfy v3 entrypoint exports every maintained non-conditional node."""
|
||||
|
||||
sys.modules.pop("prompt_control.nodes_lazy", None)
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
@@ -522,10 +170,38 @@ def test_v3_entrypoint_keeps_tile_node_when_prompt_control_unavailable(
|
||||
extension = asyncio.run(package.comfy_entrypoint())
|
||||
nodes = asyncio.run(extension.get_node_list())
|
||||
|
||||
assert [node.__name__ for node in nodes] == [
|
||||
"WD14TaggerLoaderV3",
|
||||
"TileAndTagSEGSV3",
|
||||
"SimpleLoadCheckpointV3",
|
||||
"ScaleFactorV3",
|
||||
assert _node_ids(cast(list[type[_V3Node]], nodes)) == BASE_NODE_IDS
|
||||
assert "prompt_control.nodes_lazy" not in sys.modules
|
||||
|
||||
|
||||
def test_v3_entrypoint_adds_only_prompt_control_nodes_when_available(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Prompt Control availability adds conditional nodes without removing others."""
|
||||
|
||||
sys.modules.pop("prompt_control.nodes_lazy", None)
|
||||
package = importlib.import_module("SimpleSyrup")
|
||||
nodes_v3 = importlib.import_module("SimpleSyrup.simple_syrup.nodes_v3")
|
||||
monkeypatch.setattr(nodes_v3, "prompt_control_is_available", lambda: True)
|
||||
|
||||
extension = asyncio.run(package.comfy_entrypoint())
|
||||
nodes = asyncio.run(extension.get_node_list())
|
||||
|
||||
assert _node_ids(cast(list[type[_V3Node]], nodes)) == [
|
||||
*BASE_NODE_IDS,
|
||||
*PROMPT_CONTROL_NODE_IDS,
|
||||
]
|
||||
assert "prompt_control.nodes_lazy" not in sys.modules
|
||||
|
||||
|
||||
def _node_ids(nodes: list[type[_V3Node]]) -> list[str]:
|
||||
"""Return node ids from v3 schemas."""
|
||||
|
||||
ids: list[str] = []
|
||||
for node in nodes:
|
||||
schema = node.define_schema()
|
||||
node_id: Any = schema.node_id
|
||||
if not isinstance(node_id, str):
|
||||
raise AssertionError(f"{node.__name__} has invalid node_id.")
|
||||
ids.append(node_id)
|
||||
return ids
|
||||
|
||||
@@ -0,0 +1,273 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Prompt-Control schedule and encode node schema."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from simple_syrup.nodes.schedule_and_encode_prompts_with_prompt_control import (
|
||||
ScheduleAndEncodePromptsWithPromptControl as LegacyScheduleAndEncode,
|
||||
)
|
||||
from simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control import (
|
||||
ScheduleAndEncodePromptsWithPromptControl,
|
||||
)
|
||||
|
||||
|
||||
def test_legacy_schedule_and_encode_prompt_control_node_contract() -> None:
|
||||
"""The legacy node exposes the same workflow-facing contract."""
|
||||
|
||||
inputs = LegacyScheduleAndEncode.INPUT_TYPES()
|
||||
|
||||
assert LegacyScheduleAndEncode.RETURN_TYPES == (
|
||||
"MODEL",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
)
|
||||
assert LegacyScheduleAndEncode.RETURN_NAMES == (
|
||||
"model",
|
||||
"positive",
|
||||
"negative",
|
||||
)
|
||||
assert LegacyScheduleAndEncode.FUNCTION == "execute"
|
||||
assert LegacyScheduleAndEncode.CATEGORY == "SimpleSyrup/Conditioning"
|
||||
assert list(inputs["required"]) == [
|
||||
"model",
|
||||
"clip",
|
||||
"positive_prompt",
|
||||
"negative_prompt",
|
||||
]
|
||||
assert list(inputs["optional"]) == ["encode_style"]
|
||||
assert inputs["required"]["model"][0] == "MODEL"
|
||||
assert inputs["required"]["clip"][0] == "CLIP"
|
||||
assert inputs["optional"]["encode_style"][0] == "STRING"
|
||||
assert inputs["optional"]["encode_style"][1]["forceInput"] is True
|
||||
assert inputs["required"]["positive_prompt"][1]["multiline"] is False
|
||||
assert inputs["required"]["negative_prompt"][1]["multiline"] is False
|
||||
|
||||
|
||||
def test_legacy_schedule_and_encode_prompt_control_execute_delegates(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""Legacy node execution delegates to the shared runtime builder."""
|
||||
|
||||
calls: list[dict[str, Any]] = []
|
||||
|
||||
class FakeBuilder:
|
||||
"""Runtime builder double."""
|
||||
|
||||
def build(self, **kwargs: Any) -> str:
|
||||
"""Record builder arguments and return a fixed output."""
|
||||
|
||||
calls.append(kwargs)
|
||||
return "legacy-node-output"
|
||||
|
||||
monkeypatch.setattr(
|
||||
"simple_syrup.nodes.schedule_and_encode_prompts_with_prompt_control."
|
||||
"PromptControlScheduleEncodeGraphBuilder",
|
||||
FakeBuilder,
|
||||
)
|
||||
|
||||
output = LegacyScheduleAndEncode().execute(
|
||||
model="model",
|
||||
clip="clip",
|
||||
encode_style="STYLE(A1111) ",
|
||||
positive_prompt="positive",
|
||||
negative_prompt="negative",
|
||||
)
|
||||
|
||||
assert output == "legacy-node-output"
|
||||
assert calls == [
|
||||
{
|
||||
"model": "model",
|
||||
"clip": "clip",
|
||||
"encode_style": "STYLE(A1111) ",
|
||||
"positive_prompt": "positive",
|
||||
"negative_prompt": "negative",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_legacy_schedule_and_encode_prompt_control_omits_encode_style(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""Legacy node execution no-ops style behavior when the optional socket is empty."""
|
||||
|
||||
calls: list[dict[str, Any]] = []
|
||||
|
||||
class FakeBuilder:
|
||||
"""Runtime builder double."""
|
||||
|
||||
def build(self, **kwargs: Any) -> str:
|
||||
"""Record builder arguments and return a fixed output."""
|
||||
|
||||
calls.append(kwargs)
|
||||
return "legacy-node-output"
|
||||
|
||||
monkeypatch.setattr(
|
||||
"simple_syrup.nodes.schedule_and_encode_prompts_with_prompt_control."
|
||||
"PromptControlScheduleEncodeGraphBuilder",
|
||||
FakeBuilder,
|
||||
)
|
||||
|
||||
output = LegacyScheduleAndEncode().execute(
|
||||
model="model",
|
||||
clip="clip",
|
||||
positive_prompt="positive",
|
||||
negative_prompt="negative",
|
||||
)
|
||||
|
||||
assert output == "legacy-node-output"
|
||||
assert calls == [
|
||||
{
|
||||
"model": "model",
|
||||
"clip": "clip",
|
||||
"encode_style": "",
|
||||
"positive_prompt": "positive",
|
||||
"negative_prompt": "negative",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_schedule_and_encode_prompt_control_node_schema() -> None:
|
||||
"""The v3 node exposes the planned schedule and encode contract."""
|
||||
|
||||
schema = ScheduleAndEncodePromptsWithPromptControl.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.ScheduleAndEncodePromptsWithPromptControl"
|
||||
assert schema.display_name == "Schedule & Encode Prompts"
|
||||
assert schema.enable_expand is True
|
||||
assert schema.category == "SimpleSyrup/Conditioning"
|
||||
assert [output.io_type for output in schema.outputs] == [
|
||||
"MODEL",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
]
|
||||
assert [output.id for output in schema.outputs] == [
|
||||
"model",
|
||||
"positive",
|
||||
"negative",
|
||||
]
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"model",
|
||||
"clip",
|
||||
"encode_style",
|
||||
"positive_prompt",
|
||||
"negative_prompt",
|
||||
]
|
||||
assert schema.inputs[2].optional is True
|
||||
|
||||
|
||||
def test_schedule_and_encode_prompt_control_input_types() -> None:
|
||||
"""The finalized v3 schema exposes Comfy-compatible sockets."""
|
||||
|
||||
inputs = ScheduleAndEncodePromptsWithPromptControl.INPUT_TYPES()
|
||||
|
||||
assert ScheduleAndEncodePromptsWithPromptControl.RETURN_TYPES == [
|
||||
"MODEL",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
"CONDITIONING,CONDITIONING_BATCH",
|
||||
]
|
||||
assert ScheduleAndEncodePromptsWithPromptControl.RETURN_NAMES == [
|
||||
"model",
|
||||
"positive",
|
||||
"negative",
|
||||
]
|
||||
assert list(inputs["required"]) == [
|
||||
"model",
|
||||
"clip",
|
||||
"positive_prompt",
|
||||
"negative_prompt",
|
||||
]
|
||||
assert list(inputs["optional"]) == ["encode_style"]
|
||||
assert inputs["required"]["model"][0] == "MODEL"
|
||||
assert inputs["required"]["clip"][0] == "CLIP"
|
||||
assert inputs["optional"]["encode_style"][0] == "STRING"
|
||||
assert inputs["optional"]["encode_style"][1]["forceInput"] is True
|
||||
assert inputs["required"]["positive_prompt"][1]["multiline"] is False
|
||||
assert inputs["required"]["negative_prompt"][1]["multiline"] is False
|
||||
|
||||
|
||||
def test_schedule_and_encode_prompt_control_execute_delegates(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""Node execution delegates behavior to the runtime graph builder."""
|
||||
|
||||
calls: list[dict[str, Any]] = []
|
||||
|
||||
class FakeBuilder:
|
||||
"""Runtime builder double."""
|
||||
|
||||
def build(self, **kwargs: Any) -> str:
|
||||
"""Record builder arguments and return a fixed output."""
|
||||
|
||||
calls.append(kwargs)
|
||||
return "node-output"
|
||||
|
||||
monkeypatch.setattr(
|
||||
"simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control."
|
||||
"PromptControlScheduleEncodeGraphBuilder",
|
||||
FakeBuilder,
|
||||
)
|
||||
|
||||
output = ScheduleAndEncodePromptsWithPromptControl.execute(
|
||||
model="model",
|
||||
clip="clip",
|
||||
encode_style="STYLE(A1111) ",
|
||||
positive_prompt="positive",
|
||||
negative_prompt="negative",
|
||||
)
|
||||
|
||||
assert output == "node-output"
|
||||
assert calls == [
|
||||
{
|
||||
"model": "model",
|
||||
"clip": "clip",
|
||||
"encode_style": "STYLE(A1111) ",
|
||||
"positive_prompt": "positive",
|
||||
"negative_prompt": "negative",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_schedule_and_encode_prompt_control_omits_encode_style(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""Node execution no-ops style behavior when the optional socket is empty."""
|
||||
|
||||
calls: list[dict[str, Any]] = []
|
||||
|
||||
class FakeBuilder:
|
||||
"""Runtime builder double."""
|
||||
|
||||
def build(self, **kwargs: Any) -> str:
|
||||
"""Record builder arguments and return a fixed output."""
|
||||
|
||||
calls.append(kwargs)
|
||||
return "node-output"
|
||||
|
||||
monkeypatch.setattr(
|
||||
"simple_syrup.nodes_v3.schedule_and_encode_prompts_with_prompt_control."
|
||||
"PromptControlScheduleEncodeGraphBuilder",
|
||||
FakeBuilder,
|
||||
)
|
||||
|
||||
output = ScheduleAndEncodePromptsWithPromptControl.execute(
|
||||
model="model",
|
||||
clip="clip",
|
||||
positive_prompt="positive",
|
||||
negative_prompt="negative",
|
||||
)
|
||||
|
||||
assert output == "node-output"
|
||||
assert calls == [
|
||||
{
|
||||
"model": "model",
|
||||
"clip": "clip",
|
||||
"encode_style": "",
|
||||
"positive_prompt": "positive",
|
||||
"negative_prompt": "negative",
|
||||
}
|
||||
]
|
||||
@@ -16,6 +16,7 @@ from simple_syrup.domain.segs import (
|
||||
BoundingBox,
|
||||
CropRegion,
|
||||
Segment,
|
||||
batch_segs,
|
||||
coerce_segment,
|
||||
coerce_segs,
|
||||
coerce_segs_group,
|
||||
@@ -155,6 +156,73 @@ def test_impact_segs_group_conversion_returns_list_outputs() -> None:
|
||||
assert all(isinstance(segments, list) for _header, segments in output)
|
||||
|
||||
|
||||
def test_batch_segs_preserves_input_and_segment_order() -> None:
|
||||
"""Batch SEGS flattens input payloads without reordering segments."""
|
||||
|
||||
first = (
|
||||
(16, 16),
|
||||
(
|
||||
_segment("1", CropRegion(0, 0, 2, 2), 0.9),
|
||||
_segment("2", CropRegion(2, 0, 4, 2), 0.8),
|
||||
_segment("3", CropRegion(4, 0, 6, 2), 0.7),
|
||||
),
|
||||
)
|
||||
second = (
|
||||
(16, 16),
|
||||
[
|
||||
_segment("4", CropRegion(0, 2, 2, 4), 0.6),
|
||||
_segment("5", CropRegion(2, 2, 4, 4), 0.5),
|
||||
_segment("6", CropRegion(4, 2, 6, 4), 0.4),
|
||||
],
|
||||
)
|
||||
|
||||
header, segments = batch_segs((first, second))
|
||||
|
||||
assert header == (16, 16)
|
||||
assert [segment.label for segment in segments] == ["1", "2", "3", "4", "5", "6"]
|
||||
|
||||
|
||||
def test_batch_segs_allows_empty_payloads() -> None:
|
||||
"""Empty SEGS inputs contribute no segments to the batched payload."""
|
||||
|
||||
first = (
|
||||
(16, 16),
|
||||
(
|
||||
_segment("1", CropRegion(0, 0, 2, 2), 0.9),
|
||||
_segment("2", CropRegion(2, 0, 4, 2), 0.8),
|
||||
),
|
||||
)
|
||||
empty = ((16, 16), ())
|
||||
third = ((16, 16), (_segment("3", CropRegion(4, 0, 6, 2), 0.7),))
|
||||
|
||||
_header, segments = batch_segs((first, empty, third))
|
||||
|
||||
assert [segment.label for segment in segments] == ["1", "2", "3"]
|
||||
|
||||
|
||||
def test_batch_segs_returns_empty_payload_when_all_inputs_are_empty() -> None:
|
||||
"""All-empty SEGS inputs keep the shared header and return no segments."""
|
||||
|
||||
assert batch_segs((((16, 16), ()), ((16, 16), []))) == ((16, 16), ())
|
||||
|
||||
|
||||
def test_batch_segs_rejects_no_inputs() -> None:
|
||||
"""Batch SEGS requires at least one payload for an output header."""
|
||||
|
||||
with pytest.raises(ValueError, match="one or more SEGS inputs"):
|
||||
batch_segs(())
|
||||
|
||||
|
||||
def test_batch_segs_rejects_mismatched_headers() -> None:
|
||||
"""Batch SEGS refuses to merge regions targeting different image sizes."""
|
||||
|
||||
first = ((8, 16), (_segment("first", CropRegion(0, 0, 2, 2), 0.9),))
|
||||
second = ((16, 8), (_segment("second", CropRegion(0, 0, 2, 2), 0.8),))
|
||||
|
||||
with pytest.raises(ValueError, match="input 2 is 16x8 but input 1 is 8x16"):
|
||||
batch_segs((first, second))
|
||||
|
||||
|
||||
def test_sort_order_options_are_plain_english_and_ordered() -> None:
|
||||
"""SEGS sort options match the detector node combo contract."""
|
||||
|
||||
|
||||
+74
-1
@@ -78,7 +78,12 @@ def test_saving_settings_writes_validated_schema(tmp_path: Path) -> None:
|
||||
repository.save(SimpleSyrupSettings(show_downloadable_models=False))
|
||||
|
||||
assert json.loads(path.read_text(encoding="utf-8")) == {
|
||||
"show_downloadable_models": False
|
||||
"external_llm": {
|
||||
"base_url": "",
|
||||
"cached_models": [],
|
||||
"default_model": "",
|
||||
},
|
||||
"show_downloadable_models": False,
|
||||
}
|
||||
|
||||
|
||||
@@ -100,3 +105,71 @@ def test_payload_validation_rejects_non_boolean_value() -> None:
|
||||
|
||||
with pytest.raises(SimpleSyrupSettingsError, match="show_downloadable_models"):
|
||||
SimpleSyrupSettings.from_payload({"show_downloadable_models": 1})
|
||||
|
||||
|
||||
def test_missing_external_llm_settings_loads_defaults(tmp_path: Path) -> None:
|
||||
"""Existing settings files without external LLM settings remain valid."""
|
||||
|
||||
path = tmp_path / "settings.json"
|
||||
path.write_text(
|
||||
json.dumps({"show_downloadable_models": False}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
settings = SimpleSyrupSettingsRepository(path).load()
|
||||
|
||||
assert settings.show_downloadable_models is False
|
||||
assert settings.external_llm.base_url == ""
|
||||
assert settings.external_llm.cached_models == ()
|
||||
assert settings.external_llm.default_model == ""
|
||||
|
||||
|
||||
def test_valid_external_llm_settings_are_loaded(tmp_path: Path) -> None:
|
||||
"""External LLM settings are normalized when loaded."""
|
||||
|
||||
path = tmp_path / "settings.json"
|
||||
path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"show_downloadable_models": True,
|
||||
"external_llm": {
|
||||
"base_url": "https://provider.example/v1/",
|
||||
"cached_models": ["model-a", "model-a", "model-b"],
|
||||
"default_model": "model-b",
|
||||
},
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
settings = SimpleSyrupSettingsRepository(path).load()
|
||||
|
||||
assert settings.external_llm.base_url == "https://provider.example/v1"
|
||||
assert settings.external_llm.cached_models == ("model-a", "model-b")
|
||||
assert settings.external_llm.default_model == "model-b"
|
||||
|
||||
|
||||
def test_invalid_external_llm_settings_fall_back_to_external_defaults(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Malformed external LLM settings do not invalidate other settings."""
|
||||
|
||||
path = tmp_path / "settings.json"
|
||||
path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"show_downloadable_models": False,
|
||||
"external_llm": {
|
||||
"base_url": "not-a-url",
|
||||
"cached_models": ["model-a"],
|
||||
"default_model": "model-a",
|
||||
},
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
settings = SimpleSyrupSettingsRepository(path).load()
|
||||
|
||||
assert settings.show_downloadable_models is False
|
||||
assert settings.external_llm.base_url == ""
|
||||
|
||||
@@ -17,6 +17,7 @@ from aiohttp import web
|
||||
|
||||
import simple_syrup.runtime.settings_routes as settings_routes
|
||||
from simple_syrup.runtime.settings import (
|
||||
ExternalLLMSettings,
|
||||
SimpleSyrupSettings,
|
||||
SimpleSyrupSettingsRepository,
|
||||
)
|
||||
@@ -114,7 +115,14 @@ def test_get_settings_returns_current_settings(tmp_path: Path) -> None:
|
||||
response = asyncio.run(prompt_server.routes.get_handlers[SETTINGS_ROUTE](object()))
|
||||
|
||||
assert response.status == 200
|
||||
assert json.loads(response_text(response)) == {"show_downloadable_models": False}
|
||||
assert json.loads(response_text(response)) == {
|
||||
"external_llm": {
|
||||
"base_url": "",
|
||||
"cached_models": [],
|
||||
"default_model": "",
|
||||
},
|
||||
"show_downloadable_models": False,
|
||||
}
|
||||
|
||||
|
||||
def test_post_settings_validates_and_persists_payload(tmp_path: Path) -> None:
|
||||
@@ -131,10 +139,57 @@ def test_post_settings_validates_and_persists_payload(tmp_path: Path) -> None:
|
||||
)
|
||||
|
||||
assert response.status == 200
|
||||
assert json.loads(response_text(response)) == {"show_downloadable_models": False}
|
||||
assert json.loads(response_text(response)) == {
|
||||
"external_llm": {
|
||||
"base_url": "",
|
||||
"cached_models": [],
|
||||
"default_model": "",
|
||||
},
|
||||
"show_downloadable_models": False,
|
||||
}
|
||||
assert repository.load().show_downloadable_models is False
|
||||
|
||||
|
||||
def test_post_settings_preserves_external_llm_when_payload_omits_it(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Generic settings updates do not clear saved external LLM config."""
|
||||
|
||||
prompt_server = FakePromptServer()
|
||||
repository = SimpleSyrupSettingsRepository(tmp_path / "settings.json")
|
||||
repository.save(
|
||||
SimpleSyrupSettings(
|
||||
show_downloadable_models=True,
|
||||
external_llm=ExternalLLMSettings(
|
||||
base_url="https://provider.example/v1",
|
||||
cached_models=("model-a",),
|
||||
default_model="model-a",
|
||||
),
|
||||
)
|
||||
)
|
||||
register_fake_routes(repository, prompt_server)
|
||||
|
||||
response = asyncio.run(
|
||||
prompt_server.routes.post_handlers[SETTINGS_ROUTE](
|
||||
FakeRequest({"show_downloadable_models": False})
|
||||
)
|
||||
)
|
||||
|
||||
payload = json.loads(response_text(response))
|
||||
assert response.status == 200
|
||||
assert payload == {
|
||||
"external_llm": {
|
||||
"base_url": "https://provider.example/v1",
|
||||
"cached_models": ["model-a"],
|
||||
"default_model": "model-a",
|
||||
},
|
||||
"show_downloadable_models": False,
|
||||
}
|
||||
loaded = repository.load()
|
||||
assert loaded.show_downloadable_models is False
|
||||
assert loaded.external_llm.base_url == "https://provider.example/v1"
|
||||
|
||||
|
||||
def test_post_settings_rejects_non_boolean_payload(tmp_path: Path) -> None:
|
||||
"""POST rejects malformed setting values."""
|
||||
|
||||
|
||||
@@ -0,0 +1,360 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for external-LLM tagging of existing SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, NativeSegs, Segment
|
||||
from simple_syrup.services.tag_segs_with_external_llm_service import (
|
||||
LLMTagFormattingControls,
|
||||
TagSEGSWithExternalLLMService,
|
||||
)
|
||||
|
||||
DEFAULT_CROP_REGION = CropRegion(1, 1, 3, 3)
|
||||
|
||||
|
||||
def test_service_preserves_segs_llm_prompt_and_conditioning_order(
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
"""Existing SEGS, LLM calls, formatted prompts, and conditioning stay aligned."""
|
||||
|
||||
progress = _ProgressRecorder()
|
||||
llm = _FakeLLM(("blue_hair, smile, bad_tag", "green_eyes"))
|
||||
image_encoder = _FakeImageEncoder()
|
||||
conditioning_encoder = _FakeConditioningEncoder()
|
||||
service = TagSEGSWithExternalLLMService(
|
||||
llm=llm,
|
||||
image_encoder=image_encoder,
|
||||
conditioning_encoder=conditioning_encoder,
|
||||
progress_factory=lambda _total: progress,
|
||||
)
|
||||
segs = _native_segs(("first", "second"))
|
||||
|
||||
with caplog.at_level(logging.INFO, logger="simple_syrup"):
|
||||
result = service.tag(
|
||||
image=_image(),
|
||||
segs=segs,
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="masterpiece",
|
||||
seg_image_mode="transparent mask",
|
||||
formatting=LLMTagFormattingControls(
|
||||
replace_underscore=True,
|
||||
trailing_comma=True,
|
||||
exclude_tags="bad tag",
|
||||
),
|
||||
max_tokens=64,
|
||||
reasoning_effort="off",
|
||||
)
|
||||
|
||||
assert [segment.label for segment in result.segs[1]] == ["first", "second"]
|
||||
assert image_encoder.calls == (
|
||||
("first", "transparent mask"),
|
||||
("second", "transparent mask"),
|
||||
)
|
||||
assert [call["image_data_url"] for call in llm.calls] == [
|
||||
"data:image/png;base64,first",
|
||||
"data:image/png;base64,second",
|
||||
]
|
||||
assert [call["model"] for call in llm.calls] == ["vision-model", "vision-model"]
|
||||
assert [call["max_tokens"] for call in llm.calls] == [64, 64]
|
||||
assert [call["reasoning_effort"] for call in llm.calls] == ["off", "off"]
|
||||
assert conditioning_encoder.chunks == (
|
||||
"masterpiece, blue hair, smile,",
|
||||
"masterpiece, green eyes,",
|
||||
)
|
||||
assert result.positive.entries == (
|
||||
"clip:masterpiece, blue hair, smile,",
|
||||
"clip:masterpiece, green eyes,",
|
||||
)
|
||||
assert progress.updates == [1, 1, 1, 1]
|
||||
record = caplog.records[-1]
|
||||
assert record.__dict__["operation"] == "tag_segs_with_external_llm"
|
||||
assert record.__dict__["segment_count"] == 2
|
||||
assert record.__dict__["external_llm_model"] == "vision-model"
|
||||
assert record.__dict__["seg_image_mode"] == "transparent mask"
|
||||
assert record.__dict__["universal_positive_present"] is True
|
||||
|
||||
|
||||
def test_service_preserves_underscores_when_requested() -> None:
|
||||
"""Formatting can keep booru-style underscores."""
|
||||
|
||||
conditioning_encoder = _FakeConditioningEncoder()
|
||||
service = _service(
|
||||
llm=_FakeLLM(("blue_hair, smile",)),
|
||||
conditioning_encoder=conditioning_encoder,
|
||||
)
|
||||
|
||||
service.tag(
|
||||
image=_image(),
|
||||
segs=_native_segs(("first",)),
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="",
|
||||
seg_image_mode="black mask",
|
||||
formatting=LLMTagFormattingControls(replace_underscore=False),
|
||||
)
|
||||
|
||||
assert conditioning_encoder.chunks == ("blue_hair, smile",)
|
||||
|
||||
|
||||
def test_service_rejects_empty_segs() -> None:
|
||||
"""External LLM tagging needs at least one SEG to tag."""
|
||||
|
||||
with pytest.raises(ValueError, match="No SEGS"):
|
||||
_service().tag(
|
||||
image=_image(),
|
||||
segs=((4, 4), ()),
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="",
|
||||
seg_image_mode="transparent mask",
|
||||
formatting=LLMTagFormattingControls(),
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_segs_image_header_mismatch() -> None:
|
||||
"""SEGS must describe the source image dimensions."""
|
||||
|
||||
with pytest.raises(ValueError, match="SEGS is 8x4, image is 4x4"):
|
||||
_service().tag(
|
||||
image=_image(),
|
||||
segs=((8, 4), _native_segs(("first",))[1]),
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="",
|
||||
seg_image_mode="transparent mask",
|
||||
formatting=LLMTagFormattingControls(),
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_crop_regions_outside_image() -> None:
|
||||
"""SEG crop regions must fit inside the connected source image."""
|
||||
|
||||
with pytest.raises(ValueError, match="crop_region must fit"):
|
||||
_service().tag(
|
||||
image=_image(),
|
||||
segs=_native_segs(("outside",), crop_region=CropRegion(3, 3, 5, 5)),
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="",
|
||||
seg_image_mode="transparent mask",
|
||||
formatting=LLMTagFormattingControls(),
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_empty_llm_responses() -> None:
|
||||
"""Empty provider responses cannot produce usable regional prompts."""
|
||||
|
||||
with pytest.raises(ValueError, match="empty response"):
|
||||
_service(llm=_FakeLLM((" ",))).tag(
|
||||
image=_image(),
|
||||
segs=_native_segs(("first",)),
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="",
|
||||
seg_image_mode="transparent mask",
|
||||
formatting=LLMTagFormattingControls(),
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_responses_removed_by_exclusions() -> None:
|
||||
"""Exclusion filtering must not silently create blank prompts."""
|
||||
|
||||
with pytest.raises(ValueError, match="no usable tags"):
|
||||
_service(llm=_FakeLLM(("bad_tag",))).tag(
|
||||
image=_image(),
|
||||
segs=_native_segs(("first",)),
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="",
|
||||
seg_image_mode="transparent mask",
|
||||
formatting=LLMTagFormattingControls(exclude_tags="bad tag"),
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_conditioning_count_mismatch() -> None:
|
||||
"""Conditioning output count must stay aligned to SEGS count."""
|
||||
|
||||
with pytest.raises(ValueError, match="returned 1 entries for 2 SEGS"):
|
||||
_service(
|
||||
llm=_FakeLLM(("first", "second")),
|
||||
conditioning_encoder=_ShortConditioningEncoder(),
|
||||
).tag(
|
||||
image=_image(),
|
||||
segs=_native_segs(("first", "second")),
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="",
|
||||
seg_image_mode="transparent mask",
|
||||
formatting=LLMTagFormattingControls(),
|
||||
)
|
||||
|
||||
|
||||
class _FakeLLM:
|
||||
"""Return ordered external LLM responses."""
|
||||
|
||||
def __init__(self, responses: tuple[str, ...] = ("tag",)) -> None:
|
||||
"""Store fixed responses."""
|
||||
|
||||
self.responses = responses
|
||||
self.calls: list[dict[str, object]] = []
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return deterministic model choices."""
|
||||
|
||||
return ["vision-model"]
|
||||
|
||||
def generate_with_image_data_url(
|
||||
self,
|
||||
model: str,
|
||||
system_prompt: str,
|
||||
user_prompt: str,
|
||||
max_tokens: int = 1024,
|
||||
reasoning_effort: str = "default",
|
||||
image_data_url: str | None = None,
|
||||
) -> str:
|
||||
"""Capture one LLM call and return its configured response."""
|
||||
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"system_prompt": system_prompt,
|
||||
"user_prompt": user_prompt,
|
||||
"max_tokens": max_tokens,
|
||||
"reasoning_effort": reasoning_effort,
|
||||
"image_data_url": image_data_url,
|
||||
}
|
||||
)
|
||||
return self.responses[len(self.calls) - 1]
|
||||
|
||||
|
||||
class _FakeImageEncoder:
|
||||
"""Return visible data URLs for SEG crops."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize captured calls."""
|
||||
|
||||
self.calls: tuple[tuple[str, str], ...] = ()
|
||||
|
||||
def encode_segment_as_data_url(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
segment: Segment,
|
||||
mode: str,
|
||||
) -> str:
|
||||
"""Record one SEG image encoding call."""
|
||||
|
||||
assert tuple(image.shape) == (1, 4, 4, 3)
|
||||
self.calls = (*self.calls, (segment.label, mode))
|
||||
return f"data:image/png;base64,{segment.label}"
|
||||
|
||||
|
||||
class _FakeConditioningEncoder:
|
||||
"""Return visible conditioning entries for prompt chunks."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize captured chunks."""
|
||||
|
||||
self.chunks: tuple[str, ...] = ()
|
||||
|
||||
def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch:
|
||||
"""Encode prompts as simple strings."""
|
||||
|
||||
self.chunks = chunks
|
||||
return ConditioningBatch(tuple(f"{clip}:{chunk}" for chunk in chunks))
|
||||
|
||||
|
||||
class _ShortConditioningEncoder:
|
||||
"""Return too few conditioning entries."""
|
||||
|
||||
def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch:
|
||||
"""Encode only the first prompt chunk."""
|
||||
|
||||
_ = clip
|
||||
return ConditioningBatch((chunks[0],))
|
||||
|
||||
|
||||
class _ProgressRecorder:
|
||||
"""Record service progress updates."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize captured update values."""
|
||||
|
||||
self.updates: list[int] = []
|
||||
|
||||
def update(self, value: int) -> None:
|
||||
"""Record one progress advance."""
|
||||
|
||||
self.updates.append(value)
|
||||
|
||||
|
||||
def _service(
|
||||
llm: _FakeLLM | None = None,
|
||||
conditioning_encoder: _FakeConditioningEncoder
|
||||
| _ShortConditioningEncoder
|
||||
| None = None,
|
||||
) -> TagSEGSWithExternalLLMService:
|
||||
"""Create a service with fake boundaries."""
|
||||
|
||||
return TagSEGSWithExternalLLMService(
|
||||
llm=llm or _FakeLLM(),
|
||||
image_encoder=_FakeImageEncoder(),
|
||||
conditioning_encoder=conditioning_encoder or _FakeConditioningEncoder(),
|
||||
)
|
||||
|
||||
|
||||
def _native_segs(
|
||||
labels: tuple[str, ...],
|
||||
crop_region: CropRegion = DEFAULT_CROP_REGION,
|
||||
) -> NativeSegs:
|
||||
"""Create native SEGS with stable crop regions."""
|
||||
|
||||
segments = tuple(
|
||||
Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask=torch.ones((crop_region.height, crop_region.width)),
|
||||
confidence=1.0,
|
||||
crop_region=crop_region,
|
||||
bbox=BoundingBox(
|
||||
crop_region.left,
|
||||
crop_region.top,
|
||||
crop_region.right,
|
||||
crop_region.bottom,
|
||||
),
|
||||
label=label,
|
||||
)
|
||||
for label in labels
|
||||
)
|
||||
return (4, 4), segments
|
||||
|
||||
|
||||
def _image() -> torch.Tensor:
|
||||
"""Return a small deterministic BHWC image."""
|
||||
|
||||
return torch.arange(4 * 4 * 3, dtype=torch.float32).reshape(1, 4, 4, 3) / 255.0
|
||||
@@ -0,0 +1,145 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Tag SEGS w/ External LLM Comfy v3 node."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, ImpactSegs, Segment
|
||||
from simple_syrup.nodes_v3.tag_segs_with_external_llm import (
|
||||
TagSEGSWithExternalLLMV3,
|
||||
)
|
||||
from simple_syrup.services.tag_segs_with_external_llm_service import (
|
||||
LLMTagFormattingControls,
|
||||
TagSEGSWithExternalLLMResult,
|
||||
)
|
||||
|
||||
|
||||
def test_tag_segs_with_external_llm_v3_schema(monkeypatch: Any) -> None:
|
||||
"""The v3 schema exposes the external LLM SEGS tagging contract."""
|
||||
|
||||
monkeypatch.setattr(TagSEGSWithExternalLLMV3, "_service", _FakeService())
|
||||
|
||||
schema = TagSEGSWithExternalLLMV3.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.TagSEGSWithExternalLLM"
|
||||
assert schema.display_name == "Tag SEGS w/ External LLM"
|
||||
assert schema.category == "SimpleSyrup/Detailing"
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"image",
|
||||
"segs",
|
||||
"clip",
|
||||
"model",
|
||||
"system_prompt",
|
||||
"user_prompt",
|
||||
"universal_positive",
|
||||
"seg_image_mode",
|
||||
"replace_underscore",
|
||||
"trailing_comma",
|
||||
"exclude_tags",
|
||||
"max_tokens",
|
||||
"reasoning_effort",
|
||||
]
|
||||
assert schema.inputs[0].io_type == "IMAGE"
|
||||
assert schema.inputs[1].io_type == "SEGS"
|
||||
assert schema.inputs[2].io_type == "CLIP"
|
||||
assert schema.inputs[3].options == ["vision-model"]
|
||||
assert schema.inputs[4].default == ""
|
||||
assert schema.inputs[5].default == ""
|
||||
seg_image_mode = schema.inputs[7]
|
||||
assert seg_image_mode.io_type == "COMBO"
|
||||
assert seg_image_mode.options == ["transparent mask", "black mask", "full crop"]
|
||||
assert seg_image_mode.default == "transparent mask"
|
||||
assert schema.inputs[8].default is True
|
||||
assert schema.inputs[9].default is False
|
||||
assert schema.inputs[10].default == ""
|
||||
assert schema.inputs[11].default == 1024
|
||||
assert schema.inputs[12].options == ["default", "high", "medium", "low", "off"]
|
||||
assert [output.id for output in schema.outputs] == ["segs", "positive"]
|
||||
assert [output.io_type for output in schema.outputs] == [
|
||||
"SEGS",
|
||||
"CONDITIONING_BATCH",
|
||||
]
|
||||
|
||||
|
||||
def test_tag_segs_with_external_llm_v3_execute_forwards_to_service(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""The v3 node delegates execution to the service."""
|
||||
|
||||
service = _FakeService()
|
||||
monkeypatch.setattr(TagSEGSWithExternalLLMV3, "_service", service)
|
||||
image = torch.zeros((1, 8, 8, 3))
|
||||
segs: ImpactSegs = ((8, 8), [])
|
||||
|
||||
output_segs, positive = TagSEGSWithExternalLLMV3.execute(
|
||||
image=image,
|
||||
segs=segs,
|
||||
clip="clip",
|
||||
model="vision-model",
|
||||
system_prompt="system",
|
||||
user_prompt="user",
|
||||
universal_positive="masterpiece",
|
||||
seg_image_mode="black mask",
|
||||
replace_underscore=False,
|
||||
trailing_comma=True,
|
||||
exclude_tags="bad tag",
|
||||
max_tokens=128,
|
||||
reasoning_effort="off",
|
||||
)
|
||||
|
||||
assert output_segs is service.result.segs
|
||||
assert positive is service.result.positive
|
||||
assert service.call["image"] is image
|
||||
assert service.call["segs"] is segs
|
||||
assert service.call["clip"] == "clip"
|
||||
assert service.call["model"] == "vision-model"
|
||||
assert service.call["system_prompt"] == "system"
|
||||
assert service.call["user_prompt"] == "user"
|
||||
assert service.call["universal_positive"] == "masterpiece"
|
||||
assert service.call["seg_image_mode"] == "black mask"
|
||||
assert service.call["max_tokens"] == 128
|
||||
assert service.call["reasoning_effort"] == "off"
|
||||
formatting = service.call["formatting"]
|
||||
assert isinstance(formatting, LLMTagFormattingControls)
|
||||
assert formatting.replace_underscore is False
|
||||
assert formatting.trailing_comma is True
|
||||
assert formatting.exclude_tags == "bad tag"
|
||||
|
||||
|
||||
class _FakeService:
|
||||
"""Capture v3 node service calls."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create a fake service result."""
|
||||
|
||||
segment = Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask=torch.ones((8, 8)),
|
||||
confidence=1.0,
|
||||
crop_region=CropRegion(0, 0, 8, 8),
|
||||
bbox=BoundingBox(0, 0, 8, 8),
|
||||
label="seg_001",
|
||||
)
|
||||
self.result = TagSEGSWithExternalLLMResult(
|
||||
segs=((8, 8), [segment]),
|
||||
positive=ConditioningBatch(("encoded",)),
|
||||
)
|
||||
self.call: dict[str, object] = {}
|
||||
|
||||
def model_choices(self) -> list[str]:
|
||||
"""Return deterministic model choices."""
|
||||
|
||||
return ["vision-model"]
|
||||
|
||||
def tag(self, **kwargs: object) -> TagSEGSWithExternalLLMResult:
|
||||
"""Return a fixed result and remember provided inputs."""
|
||||
|
||||
self.call = kwargs
|
||||
return self.result
|
||||
@@ -0,0 +1,112 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Tag SEGS w/ WD14 node contract."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, ImpactSegs, Segment
|
||||
from simple_syrup.nodes.tag_segs_with_wd14 import TagSEGSWithWD14
|
||||
from simple_syrup.nodes.tile_and_tag_segs import DEFAULT_EXCLUDE_TAGS
|
||||
from simple_syrup.runtime.wd14_tagger import WD14TagFormattingControls
|
||||
from simple_syrup.services.tag_segs_with_wd14_service import TagSEGSWithWD14Result
|
||||
|
||||
|
||||
def test_tag_segs_with_wd14_contract() -> None:
|
||||
"""Tag SEGS w/ WD14 exposes the agreed ComfyUI contract."""
|
||||
|
||||
inputs = TagSEGSWithWD14.INPUT_TYPES()
|
||||
|
||||
assert TagSEGSWithWD14.RETURN_TYPES == ("SEGS", "CONDITIONING_BATCH")
|
||||
assert TagSEGSWithWD14.RETURN_NAMES == ("segs", "positive")
|
||||
assert TagSEGSWithWD14.FUNCTION == "tag"
|
||||
assert TagSEGSWithWD14.CATEGORY == "SimpleSyrup/Detailing"
|
||||
assert list(inputs["required"]) == [
|
||||
"image",
|
||||
"segs",
|
||||
"clip",
|
||||
"wd14_tagger",
|
||||
"universal_positive",
|
||||
"threshold",
|
||||
"character_threshold",
|
||||
"replace_underscore",
|
||||
"trailing_comma",
|
||||
"exclude_tags",
|
||||
]
|
||||
assert inputs["required"]["segs"][0] == "SEGS"
|
||||
assert inputs["required"]["clip"][0] == "CLIP"
|
||||
assert inputs["required"]["wd14_tagger"][0] == "WD14_TAGGER"
|
||||
assert inputs["required"]["universal_positive"][0] == "STRING"
|
||||
assert inputs["required"]["universal_positive"][1]["default"] == ""
|
||||
assert inputs["required"]["threshold"][1]["default"] == 0.35
|
||||
assert inputs["required"]["character_threshold"][1]["default"] == 1.0
|
||||
assert inputs["required"]["replace_underscore"][1]["default"] is True
|
||||
assert inputs["required"]["trailing_comma"][1]["default"] is False
|
||||
assert inputs["required"]["exclude_tags"][1]["default"] == DEFAULT_EXCLUDE_TAGS
|
||||
assert "optional" not in inputs
|
||||
|
||||
|
||||
def test_tag_segs_with_wd14_delegates_to_service(monkeypatch: Any) -> None:
|
||||
"""The node delegates behavior and returns service outputs unchanged."""
|
||||
|
||||
service = _FakeService()
|
||||
monkeypatch.setattr(TagSEGSWithWD14, "service_class", lambda: service)
|
||||
image = torch.zeros((1, 8, 8, 3))
|
||||
segs: ImpactSegs = ((8, 8), [])
|
||||
wd14_tagger = object()
|
||||
|
||||
output_segs, positive = TagSEGSWithWD14().tag(
|
||||
image=image,
|
||||
segs=segs,
|
||||
clip="clip",
|
||||
wd14_tagger=wd14_tagger,
|
||||
universal_positive="masterpiece",
|
||||
threshold=0.35,
|
||||
character_threshold=1.0,
|
||||
replace_underscore=True,
|
||||
trailing_comma=False,
|
||||
exclude_tags=DEFAULT_EXCLUDE_TAGS,
|
||||
)
|
||||
|
||||
assert output_segs is service.result.segs
|
||||
assert positive is service.result.positive
|
||||
assert service.call["image"] is image
|
||||
assert service.call["segs"] is segs
|
||||
assert service.call["clip"] == "clip"
|
||||
assert service.call["wd14_tagger"] is wd14_tagger
|
||||
assert service.call["universal_positive"] == "masterpiece"
|
||||
tag_controls = cast(WD14TagFormattingControls, service.call["tag_controls"])
|
||||
assert tag_controls.threshold == 0.35
|
||||
|
||||
|
||||
class _FakeService:
|
||||
"""Capture node calls for delegation tests."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create a fake service result."""
|
||||
|
||||
segment = Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask=torch.ones((8, 8)),
|
||||
confidence=1.0,
|
||||
crop_region=CropRegion(0, 0, 8, 8),
|
||||
bbox=BoundingBox(0, 0, 8, 8),
|
||||
label="seg_001",
|
||||
)
|
||||
self.result = TagSEGSWithWD14Result(
|
||||
segs=((8, 8), [segment]),
|
||||
positive=ConditioningBatch(("encoded",)),
|
||||
)
|
||||
self.call: dict[str, object] = {}
|
||||
|
||||
def tag(self, **kwargs: object) -> TagSEGSWithWD14Result:
|
||||
"""Return a fixed result and remember provided inputs."""
|
||||
|
||||
self.call = kwargs
|
||||
return self.result
|
||||
@@ -0,0 +1,283 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for WD14 tagging of existing SEGS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, NativeSegs, Segment
|
||||
from simple_syrup.runtime.loaded_models import LoadedWD14Tagger
|
||||
from simple_syrup.runtime.wd14_tagger import (
|
||||
FloatArray,
|
||||
WD14TagFormattingControls,
|
||||
WD14TagRecord,
|
||||
)
|
||||
from simple_syrup.services.tag_segs_with_wd14_service import TagSEGSWithWD14Service
|
||||
|
||||
|
||||
def test_service_preserves_existing_segs_tag_and_conditioning_order() -> None:
|
||||
"""Existing SEGS, crops, tags, and conditioning stay aligned by index."""
|
||||
|
||||
progress = _ProgressRecorder()
|
||||
tagger = _FakeTagger(("tag first", "", "tag third"))
|
||||
encoder = _FakeEncoder()
|
||||
loaded_tagger = _loaded_tagger()
|
||||
service = TagSEGSWithWD14Service(
|
||||
tagger=tagger,
|
||||
encoder=encoder,
|
||||
progress_factory=lambda _total: progress,
|
||||
)
|
||||
segs = _native_segs(("first", "second", "third"))
|
||||
|
||||
result = service.tag(
|
||||
image=_image(),
|
||||
segs=segs,
|
||||
clip="clip",
|
||||
wd14_tagger=loaded_tagger,
|
||||
tag_controls=_tag_controls(),
|
||||
universal_positive="masterpiece",
|
||||
)
|
||||
|
||||
assert [segment.label for segment in result.segs[1]] == [
|
||||
"first",
|
||||
"second",
|
||||
"third",
|
||||
]
|
||||
assert [tuple(crop.shape) for crop in tagger.crops] == [
|
||||
(1, 2, 2, 3),
|
||||
(1, 2, 2, 3),
|
||||
(1, 2, 2, 3),
|
||||
]
|
||||
assert tagger.loaded_tagger is loaded_tagger
|
||||
assert encoder.chunks == (
|
||||
"masterpiece, tag first",
|
||||
"masterpiece",
|
||||
"masterpiece, tag third",
|
||||
)
|
||||
assert result.positive.entries == (
|
||||
"clip:masterpiece, tag first",
|
||||
"clip:masterpiece",
|
||||
"clip:masterpiece, tag third",
|
||||
)
|
||||
assert progress.updates == [1, 3, 1]
|
||||
|
||||
|
||||
def test_service_rejects_empty_segs() -> None:
|
||||
"""Tagging empty SEGS would not produce a selectable conditioning batch."""
|
||||
|
||||
service = TagSEGSWithWD14Service(
|
||||
tagger=_FakeTagger(()),
|
||||
encoder=_FakeEncoder(),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="No SEGS"):
|
||||
service.tag(
|
||||
image=_image(),
|
||||
segs=((4, 4), ()),
|
||||
clip="clip",
|
||||
wd14_tagger=_loaded_tagger(),
|
||||
tag_controls=_tag_controls(),
|
||||
universal_positive="",
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_segs_image_header_mismatch() -> None:
|
||||
"""SEGS must describe the image being cropped for tagging."""
|
||||
|
||||
service = TagSEGSWithWD14Service(
|
||||
tagger=_FakeTagger(("tag",)),
|
||||
encoder=_FakeEncoder(),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="SEGS is 8x4, image is 4x4"):
|
||||
service.tag(
|
||||
image=_image(),
|
||||
segs=((8, 4), _native_segs(("first",))[1]),
|
||||
clip="clip",
|
||||
wd14_tagger=_loaded_tagger(),
|
||||
tag_controls=_tag_controls(),
|
||||
universal_positive="",
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_tagger_count_mismatch() -> None:
|
||||
"""Dropping a tag would break SEGS alignment and is rejected."""
|
||||
|
||||
service = TagSEGSWithWD14Service(
|
||||
tagger=_FakeTagger(("only one",)),
|
||||
encoder=_FakeEncoder(),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="returned 1 tag"):
|
||||
service.tag(
|
||||
image=_image(),
|
||||
segs=_native_segs(("first", "second")),
|
||||
clip="clip",
|
||||
wd14_tagger=_loaded_tagger(),
|
||||
tag_controls=_tag_controls(),
|
||||
universal_positive="",
|
||||
)
|
||||
|
||||
|
||||
def test_service_rejects_conditioning_count_mismatch() -> None:
|
||||
"""Dropping encoded conditioning would break SEGS alignment and is rejected."""
|
||||
|
||||
service = TagSEGSWithWD14Service(
|
||||
tagger=_FakeTagger(("first", "second")),
|
||||
encoder=_ShortEncoder(),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="returned 1 entries for 2 SEGS"):
|
||||
service.tag(
|
||||
image=_image(),
|
||||
segs=_native_segs(("first", "second")),
|
||||
clip="clip",
|
||||
wd14_tagger=_loaded_tagger(),
|
||||
tag_controls=_tag_controls(),
|
||||
universal_positive="",
|
||||
)
|
||||
|
||||
|
||||
class _FakeTagger:
|
||||
"""Return fixed tag strings for ordered crops."""
|
||||
|
||||
def __init__(self, tags: tuple[str, ...]) -> None:
|
||||
"""Store the fixed tags."""
|
||||
|
||||
self.tags = tags
|
||||
self.crops: tuple[torch.Tensor, ...] = ()
|
||||
self.loaded_tagger: LoadedWD14Tagger | None = None
|
||||
|
||||
def tag_images(
|
||||
self,
|
||||
loaded_tagger: LoadedWD14Tagger,
|
||||
images: tuple[torch.Tensor, ...],
|
||||
controls: WD14TagFormattingControls,
|
||||
progress: object | None = None,
|
||||
) -> tuple[str, ...]:
|
||||
"""Return fixed tags and remember the crop order."""
|
||||
|
||||
_ = controls
|
||||
if progress is not None:
|
||||
progress.update(len(images)) # type: ignore[attr-defined]
|
||||
self.loaded_tagger = loaded_tagger
|
||||
self.crops = images
|
||||
return self.tags
|
||||
|
||||
|
||||
class _FakeEncoder:
|
||||
"""Return visible conditioning values for prompt chunks."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize captured chunks."""
|
||||
|
||||
self.chunks: tuple[str, ...] = ()
|
||||
|
||||
def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch:
|
||||
"""Encode prompts as simple strings."""
|
||||
|
||||
self.chunks = chunks
|
||||
return ConditioningBatch(tuple(f"{clip}:{chunk}" for chunk in chunks))
|
||||
|
||||
|
||||
class _ShortEncoder:
|
||||
"""Return too few conditioning entries for validation tests."""
|
||||
|
||||
def encode_batch(self, clip: Any, chunks: tuple[str, ...]) -> ConditioningBatch:
|
||||
"""Encode only the first prompt chunk."""
|
||||
|
||||
_ = clip
|
||||
return ConditioningBatch((chunks[0],))
|
||||
|
||||
|
||||
class _ProgressRecorder:
|
||||
"""Record service progress updates."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize captured update values."""
|
||||
|
||||
self.updates: list[int] = []
|
||||
|
||||
def update(self, value: int) -> None:
|
||||
"""Record one progress advance."""
|
||||
|
||||
self.updates.append(value)
|
||||
|
||||
|
||||
def _native_segs(labels: tuple[str, ...]) -> NativeSegs:
|
||||
"""Create native SEGS with stable two-pixel crop regions."""
|
||||
|
||||
segments = tuple(
|
||||
Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask=torch.ones((2, 2)),
|
||||
confidence=1.0,
|
||||
crop_region=CropRegion(index, index, index + 2, index + 2),
|
||||
bbox=BoundingBox(index, index, index + 2, index + 2),
|
||||
label=label,
|
||||
)
|
||||
for index, label in enumerate(labels)
|
||||
)
|
||||
return (4, 4), segments
|
||||
|
||||
|
||||
def _image() -> torch.Tensor:
|
||||
"""Return a small deterministic BHWC image."""
|
||||
|
||||
return torch.arange(4 * 4 * 3, dtype=torch.float32).reshape(1, 4, 4, 3) / 255.0
|
||||
|
||||
|
||||
def _tag_controls() -> WD14TagFormattingControls:
|
||||
"""Return valid WD14 controls for service tests."""
|
||||
|
||||
return WD14TagFormattingControls(
|
||||
threshold=0.35,
|
||||
character_threshold=1.0,
|
||||
replace_underscore=True,
|
||||
trailing_comma=False,
|
||||
exclude_tags="",
|
||||
)
|
||||
|
||||
|
||||
def _loaded_tagger() -> LoadedWD14Tagger:
|
||||
"""Return a reusable loaded WD14 tagger test container."""
|
||||
|
||||
return LoadedWD14Tagger(
|
||||
model_id="wd-eva02-large-tagger-v3",
|
||||
source="test",
|
||||
onnx_path=Path("wd-eva02-large-tagger-v3.onnx"),
|
||||
csv_path=Path("wd-eva02-large-tagger-v3.csv"),
|
||||
providers=("CPUExecutionProvider",),
|
||||
session=_FakeWD14Session(),
|
||||
tags=(WD14TagRecord("blue_hair", "0"),),
|
||||
)
|
||||
|
||||
|
||||
class _FakeWD14Session:
|
||||
"""Minimal WD14 session test double."""
|
||||
|
||||
def get_inputs(self) -> list[object]:
|
||||
"""Return no fake inputs."""
|
||||
|
||||
return []
|
||||
|
||||
def get_outputs(self) -> list[object]:
|
||||
"""Return no fake outputs."""
|
||||
|
||||
return []
|
||||
|
||||
def run(
|
||||
self, output_names: list[str], feeds: dict[str, FloatArray]
|
||||
) -> list[object]:
|
||||
"""Return no fake outputs."""
|
||||
|
||||
_ = output_names, feeds
|
||||
return []
|
||||
@@ -0,0 +1,103 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the Tag SEGS w/ WD14 Comfy v3 wrapper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.domain.segs import BoundingBox, CropRegion, ImpactSegs, Segment
|
||||
from simple_syrup.nodes.tag_segs_with_wd14 import TagSEGSWithWD14
|
||||
from simple_syrup.nodes_v3.tag_segs_with_wd14 import TagSEGSWithWD14V3
|
||||
from simple_syrup.services.tag_segs_with_wd14_service import TagSEGSWithWD14Result
|
||||
|
||||
|
||||
def test_tag_segs_with_wd14_v3_schema_includes_clip_and_wd14_tagger() -> None:
|
||||
"""The v3 schema exposes existing-SEGS WD14 tagging inputs."""
|
||||
|
||||
schema = TagSEGSWithWD14V3.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.TagSEGSWithWD14"
|
||||
assert schema.display_name == "Tag SEGS w/ WD14"
|
||||
assert [input_item.id for input_item in schema.inputs][:4] == [
|
||||
"image",
|
||||
"segs",
|
||||
"clip",
|
||||
"wd14_tagger",
|
||||
]
|
||||
assert schema.inputs[1].io_type == "SEGS"
|
||||
assert schema.inputs[2].io_type == "CLIP"
|
||||
assert schema.inputs[3].io_type == "WD14_TAGGER"
|
||||
universal_positive = schema.inputs[4]
|
||||
assert universal_positive.io_type == "STRING"
|
||||
assert universal_positive.default == ""
|
||||
assert universal_positive.multiline is False
|
||||
assert [output.id for output in schema.outputs] == ["segs", "positive"]
|
||||
assert [output.io_type for output in schema.outputs] == [
|
||||
"SEGS",
|
||||
"CONDITIONING_BATCH",
|
||||
]
|
||||
|
||||
|
||||
def test_tag_segs_with_wd14_v3_execute_forwards_to_legacy_node(
|
||||
monkeypatch: Any,
|
||||
) -> None:
|
||||
"""The v3 wrapper forwards execution to the legacy implementation."""
|
||||
|
||||
service = _FakeService()
|
||||
monkeypatch.setattr(TagSEGSWithWD14, "service_class", lambda: service)
|
||||
image = torch.zeros((1, 8, 8, 3))
|
||||
segs: ImpactSegs = ((8, 8), [])
|
||||
wd14_tagger = object()
|
||||
|
||||
output_segs, positive = TagSEGSWithWD14V3.execute(
|
||||
image=image,
|
||||
segs=segs,
|
||||
clip="clip",
|
||||
wd14_tagger=wd14_tagger,
|
||||
universal_positive="masterpiece",
|
||||
threshold=0.35,
|
||||
character_threshold=1.0,
|
||||
replace_underscore=True,
|
||||
trailing_comma=False,
|
||||
exclude_tags="",
|
||||
)
|
||||
|
||||
assert output_segs is service.result.segs
|
||||
assert positive is service.result.positive
|
||||
assert service.call["segs"] is segs
|
||||
assert service.call["clip"] == "clip"
|
||||
assert service.call["wd14_tagger"] is wd14_tagger
|
||||
assert service.call["universal_positive"] == "masterpiece"
|
||||
|
||||
|
||||
class _FakeService:
|
||||
"""Capture v3 wrapper calls through the legacy node."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Create a fake service result."""
|
||||
|
||||
segment = Segment(
|
||||
cropped_image=None,
|
||||
cropped_mask=torch.ones((8, 8)),
|
||||
confidence=1.0,
|
||||
crop_region=CropRegion(0, 0, 8, 8),
|
||||
bbox=BoundingBox(0, 0, 8, 8),
|
||||
label="seg_001",
|
||||
)
|
||||
self.result = TagSEGSWithWD14Result(
|
||||
segs=((8, 8), [segment]),
|
||||
positive=ConditioningBatch(("encoded",)),
|
||||
)
|
||||
self.call: dict[str, object] = {}
|
||||
|
||||
def tag(self, **kwargs: object) -> TagSEGSWithWD14Result:
|
||||
"""Return a fixed result and remember provided inputs."""
|
||||
|
||||
self.call = kwargs
|
||||
return self.result
|
||||
@@ -50,7 +50,7 @@ def test_validate_tiled_diffusion_mode_rejects_unsupported_mode() -> None:
|
||||
|
||||
|
||||
def test_plan_clamps_tile_size_to_latent_dimensions() -> None:
|
||||
"""Requested tiles larger than the latent are clamped to latent dimensions."""
|
||||
"""Requested tiles and overlap larger than the latent use safe dimensions."""
|
||||
|
||||
plan = build_tiled_diffusion_plan(
|
||||
latent_width=12,
|
||||
@@ -63,13 +63,13 @@ def test_plan_clamps_tile_size_to_latent_dimensions() -> None:
|
||||
|
||||
assert plan.tile_width == 12
|
||||
assert plan.tile_height == 8
|
||||
assert plan.overlap == 48
|
||||
assert plan.overlap == 4
|
||||
assert plan.tiles == (LatentTile(0, 0, 12, 8),)
|
||||
assert not tile_is_splittable(12, 8, 96, 96, 48)
|
||||
|
||||
|
||||
def test_plan_clamps_overlap_against_requested_tile_dimensions() -> None:
|
||||
"""Overlap clamps against requested tile dimensions."""
|
||||
"""Overlap clamps against effective tile dimensions."""
|
||||
|
||||
plan = build_tiled_diffusion_plan(
|
||||
latent_width=8,
|
||||
@@ -82,7 +82,26 @@ def test_plan_clamps_overlap_against_requested_tile_dimensions() -> None:
|
||||
|
||||
assert plan.tile_width == 8
|
||||
assert plan.tile_height == 8
|
||||
assert plan.overlap == 92
|
||||
assert plan.overlap == 4
|
||||
|
||||
|
||||
def test_plan_uses_single_tile_for_small_latent_with_default_detailer_overlap() -> None:
|
||||
"""Small detailer latents use one full tile instead of a zero-stride grid."""
|
||||
|
||||
plan = build_tiled_diffusion_plan(
|
||||
latent_width=16,
|
||||
latent_height=30,
|
||||
tile_width=128,
|
||||
tile_height=128,
|
||||
overlap=16,
|
||||
tile_batch_size=4,
|
||||
)
|
||||
|
||||
assert plan.tile_width == 16
|
||||
assert plan.tile_height == 30
|
||||
assert plan.overlap == 12
|
||||
assert plan.tiles == (LatentTile(0, 0, 16, 30),)
|
||||
assert not tile_is_splittable(16, 30, 128, 128, 16)
|
||||
|
||||
|
||||
def test_plan_generates_row_major_symmetric_tiles() -> None:
|
||||
|
||||
@@ -11,6 +11,7 @@ from typing import Any
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from simple_syrup.domain.conditioning_batch import ConditioningBatch
|
||||
from simple_syrup.services.tiled_diffusion_sampling_service import (
|
||||
TiledDiffusionSamplingService,
|
||||
)
|
||||
@@ -129,6 +130,92 @@ def test_service_forwards_sampling_arguments_unchanged(
|
||||
}
|
||||
|
||||
|
||||
def test_service_forwards_differential_diffusion_request(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""The dispatcher preserves differential-denoise-mask composition requests."""
|
||||
|
||||
calls: dict[str, Any] = {}
|
||||
|
||||
def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]:
|
||||
"""Record forwarded arguments."""
|
||||
|
||||
calls.update(kwargs)
|
||||
return {"samples": torch.ones((1, 4, 4, 4))}
|
||||
|
||||
monkeypatch.setattr(
|
||||
"simple_syrup.services.tiled_diffusion_sampling_service."
|
||||
"multidiffusion_sampling.sample_multidiffusion",
|
||||
fake_multidiffusion,
|
||||
)
|
||||
|
||||
TiledDiffusionSamplingService().sample(
|
||||
**(
|
||||
_sample_kwargs(diffusion_mode="multidiffusion")
|
||||
| {"differential_diffusion": True}
|
||||
)
|
||||
)
|
||||
|
||||
assert calls["differential_diffusion"] is True
|
||||
|
||||
|
||||
def test_service_selects_conditioning_batch_per_latent_item(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Batch conditioning is selected before tiled runtime dispatch."""
|
||||
|
||||
calls: list[dict[str, Any]] = []
|
||||
|
||||
def fake_multidiffusion(**kwargs: Any) -> dict[str, Any]:
|
||||
"""Record per-item runtime arguments and return marked samples."""
|
||||
|
||||
calls.append(kwargs)
|
||||
return {
|
||||
"samples": torch.full_like(
|
||||
kwargs["latent_image"]["samples"],
|
||||
float(len(calls)),
|
||||
)
|
||||
}
|
||||
|
||||
monkeypatch.setattr(
|
||||
"simple_syrup.services.tiled_diffusion_sampling_service."
|
||||
"multidiffusion_sampling.sample_multidiffusion",
|
||||
fake_multidiffusion,
|
||||
)
|
||||
|
||||
latent_samples = torch.zeros((2, 4, 4, 4))
|
||||
noise_mask = torch.ones((2, 1, 4, 4))
|
||||
result = TiledDiffusionSamplingService().sample(
|
||||
**(
|
||||
_sample_kwargs(diffusion_mode="multidiffusion")
|
||||
| {
|
||||
"positive": ConditioningBatch(("positive-0", "positive-1")),
|
||||
"negative": ConditioningBatch(("negative-last",)),
|
||||
"latent_image": {
|
||||
"samples": latent_samples,
|
||||
"batch_index": [7, 11],
|
||||
"noise_mask": noise_mask,
|
||||
"downscale_ratio_spacial": 2,
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert len(calls) == 2
|
||||
assert calls[0]["positive"] == "positive-0"
|
||||
assert calls[1]["positive"] == "positive-1"
|
||||
assert calls[0]["negative"] == "negative-last"
|
||||
assert calls[1]["negative"] == "negative-last"
|
||||
assert calls[0]["latent_image"]["batch_index"] == [7]
|
||||
assert calls[1]["latent_image"]["batch_index"] == [11]
|
||||
assert torch.equal(calls[0]["latent_image"]["noise_mask"], noise_mask[0:1])
|
||||
assert torch.equal(calls[1]["latent_image"]["noise_mask"], noise_mask[1:2])
|
||||
assert "downscale_ratio_spacial" not in result
|
||||
assert result["samples"].shape == latent_samples.shape
|
||||
assert torch.equal(result["samples"][0], torch.full((4, 4, 4), 1.0))
|
||||
assert torch.equal(result["samples"][1], torch.full((4, 4, 4), 2.0))
|
||||
|
||||
|
||||
def test_invalid_mode_fails_before_runtime_call(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
@@ -181,4 +268,5 @@ def _sample_kwargs(
|
||||
"latent_tile_overlap": 24,
|
||||
"latent_tile_batch_size": 3,
|
||||
"preview_context": preview_context,
|
||||
"differential_diffusion": False,
|
||||
}
|
||||
|
||||
@@ -54,7 +54,8 @@ def test_validate_latent_samples_rejects_non_tensor() -> None:
|
||||
def test_validate_tensor_shape_rejects_nested_tensor() -> None:
|
||||
"""Nested tensors are rejected before spatial tiling."""
|
||||
|
||||
samples = torch.nested.nested_tensor([torch.zeros((4, 8, 8))])
|
||||
with pytest.warns(UserWarning, match="nested tensors.*prototype stage"):
|
||||
samples = torch.nested.nested_tensor([torch.zeros((4, 8, 8))])
|
||||
|
||||
with pytest.raises(ValueError, match="non-nested latent samples"):
|
||||
tiled_sampling.validate_tensor_shape(samples, sampler_label="TestSampler")
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the VAE Decode (Options) Comfy v3 wrapper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
from simple_syrup.nodes import tooltips
|
||||
from simple_syrup.nodes.vae_options import (
|
||||
VAE_DECODE_TILE_SIZE_STEP,
|
||||
VAE_OVERLAP_DEFAULT,
|
||||
VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
VAE_TILE_SIZE_DEFAULT,
|
||||
VAEDecodeOptions,
|
||||
)
|
||||
from simple_syrup.nodes_v3.vae_decode_options import VAEDecodeOptionsV3
|
||||
|
||||
|
||||
def test_vae_decode_options_v3_schema() -> None:
|
||||
"""VAE Decode (Options) v3 schema mirrors the legacy node contract."""
|
||||
|
||||
schema = VAEDecodeOptionsV3.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.VAEDecodeOptions"
|
||||
assert schema.display_name == "VAE Decode (Options)"
|
||||
assert schema.category == "SimpleSyrup/Latent"
|
||||
assert schema.description == VAEDecodeOptions.DESCRIPTION
|
||||
assert schema.enable_expand is True
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"use_tiling",
|
||||
"samples",
|
||||
"vae",
|
||||
"tile_size",
|
||||
"overlap",
|
||||
"temporal_size",
|
||||
"temporal_overlap",
|
||||
]
|
||||
assert schema.inputs[0].io_type == "BOOLEAN"
|
||||
assert schema.inputs[0].default is False
|
||||
assert schema.inputs[0].tooltip == tooltips.VAE_OPTIONS_USE_TILING
|
||||
assert schema.inputs[1].io_type == "LATENT"
|
||||
assert schema.inputs[1].rawLink is True
|
||||
assert schema.inputs[2].io_type == "VAE"
|
||||
assert schema.inputs[2].rawLink is True
|
||||
assert schema.inputs[3].default == VAE_TILE_SIZE_DEFAULT
|
||||
assert schema.inputs[3].step == VAE_DECODE_TILE_SIZE_STEP
|
||||
assert schema.inputs[4].default == VAE_OVERLAP_DEFAULT
|
||||
assert schema.inputs[5].default == VAE_TEMPORAL_SIZE_DEFAULT
|
||||
assert schema.inputs[6].default == VAE_TEMPORAL_OVERLAP_DEFAULT
|
||||
assert [output.id for output in schema.outputs] == ["image"]
|
||||
assert schema.outputs[0].io_type == "IMAGE"
|
||||
assert schema.outputs[0].tooltip == tooltips.VAE_OPTIONS_IMAGE_OUTPUT
|
||||
|
||||
|
||||
def test_vae_decode_options_v3_execute_delegates_to_legacy_node() -> None:
|
||||
"""VAE Decode (Options) v3 execution returns legacy expansion behavior."""
|
||||
|
||||
graph_utils = import_module("comfy_execution.graph_utils")
|
||||
graph_utils.GraphBuilder.set_default_prefix("V3_DECODE", 0, 0)
|
||||
|
||||
result = VAEDecodeOptionsV3.execute(
|
||||
True,
|
||||
["latent", 0],
|
||||
["loader", 2],
|
||||
1024,
|
||||
128,
|
||||
96,
|
||||
16,
|
||||
)
|
||||
node = next(iter(result["expand"].values()))
|
||||
|
||||
assert node["class_type"] == "VAEDecodeTiled"
|
||||
assert node["inputs"]["tile_size"] == 1024
|
||||
assert result["result"][0][1] == 0
|
||||
@@ -0,0 +1,78 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for the VAE Encode (Options) Comfy v3 wrapper."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
|
||||
from simple_syrup.nodes import tooltips
|
||||
from simple_syrup.nodes.vae_options import (
|
||||
VAE_ENCODE_TILE_SIZE_STEP,
|
||||
VAE_OVERLAP_DEFAULT,
|
||||
VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
VAE_TILE_SIZE_DEFAULT,
|
||||
VAEEncodeOptions,
|
||||
)
|
||||
from simple_syrup.nodes_v3.vae_encode_options import VAEEncodeOptionsV3
|
||||
|
||||
|
||||
def test_vae_encode_options_v3_schema() -> None:
|
||||
"""VAE Encode (Options) v3 schema mirrors the legacy node contract."""
|
||||
|
||||
schema = VAEEncodeOptionsV3.define_schema()
|
||||
|
||||
assert schema.node_id == "SimpleSyrup.VAEEncodeOptions"
|
||||
assert schema.display_name == "VAE Encode (Options)"
|
||||
assert schema.category == "SimpleSyrup/Latent"
|
||||
assert schema.description == VAEEncodeOptions.DESCRIPTION
|
||||
assert schema.enable_expand is True
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"use_tiling",
|
||||
"pixels",
|
||||
"vae",
|
||||
"tile_size",
|
||||
"overlap",
|
||||
"temporal_size",
|
||||
"temporal_overlap",
|
||||
]
|
||||
assert schema.inputs[0].io_type == "BOOLEAN"
|
||||
assert schema.inputs[0].default is False
|
||||
assert schema.inputs[0].tooltip == tooltips.VAE_OPTIONS_USE_TILING
|
||||
assert schema.inputs[1].io_type == "IMAGE"
|
||||
assert schema.inputs[1].rawLink is True
|
||||
assert schema.inputs[2].io_type == "VAE"
|
||||
assert schema.inputs[2].rawLink is True
|
||||
assert schema.inputs[3].default == VAE_TILE_SIZE_DEFAULT
|
||||
assert schema.inputs[3].step == VAE_ENCODE_TILE_SIZE_STEP
|
||||
assert schema.inputs[4].default == VAE_OVERLAP_DEFAULT
|
||||
assert schema.inputs[5].default == VAE_TEMPORAL_SIZE_DEFAULT
|
||||
assert schema.inputs[6].default == VAE_TEMPORAL_OVERLAP_DEFAULT
|
||||
assert [output.id for output in schema.outputs] == ["latent"]
|
||||
assert schema.outputs[0].io_type == "LATENT"
|
||||
assert schema.outputs[0].tooltip == tooltips.VAE_OPTIONS_LATENT_OUTPUT
|
||||
|
||||
|
||||
def test_vae_encode_options_v3_execute_delegates_to_legacy_node() -> None:
|
||||
"""VAE Encode (Options) v3 execution returns legacy expansion behavior."""
|
||||
|
||||
graph_utils = import_module("comfy_execution.graph_utils")
|
||||
graph_utils.GraphBuilder.set_default_prefix("V3_ENCODE", 0, 0)
|
||||
|
||||
result = VAEEncodeOptionsV3.execute(
|
||||
True,
|
||||
["image", 0],
|
||||
["loader", 2],
|
||||
768,
|
||||
96,
|
||||
48,
|
||||
12,
|
||||
)
|
||||
node = next(iter(result["expand"].values()))
|
||||
|
||||
assert node["class_type"] == "VAEEncodeTiled"
|
||||
assert node["inputs"]["tile_size"] == 768
|
||||
assert result["result"][0][1] == 0
|
||||
@@ -0,0 +1,226 @@
|
||||
# SimpleSyrup - workflow-focused ComfyUI extensions for image generation
|
||||
# Copyright (C) 2026 Artificial Sweetener and contributors
|
||||
# SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
"""Tests for VAE Encode/Decode Options nodes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
from simple_syrup.nodes.vae_options import (
|
||||
VAE_DECODE_TILE_SIZE_STEP,
|
||||
VAE_ENCODE_TILE_SIZE_STEP,
|
||||
VAE_OVERLAP_DEFAULT,
|
||||
VAE_OVERLAP_MAX,
|
||||
VAE_OVERLAP_MIN,
|
||||
VAE_OVERLAP_STEP,
|
||||
VAE_TEMPORAL_OVERLAP_DEFAULT,
|
||||
VAE_TEMPORAL_OVERLAP_MAX,
|
||||
VAE_TEMPORAL_OVERLAP_MIN,
|
||||
VAE_TEMPORAL_OVERLAP_STEP,
|
||||
VAE_TEMPORAL_SIZE_DEFAULT,
|
||||
VAE_TEMPORAL_SIZE_MAX,
|
||||
VAE_TEMPORAL_SIZE_MIN,
|
||||
VAE_TEMPORAL_SIZE_STEP,
|
||||
VAE_TILE_SIZE_DEFAULT,
|
||||
VAE_TILE_SIZE_MAX,
|
||||
VAE_TILE_SIZE_MIN,
|
||||
VAEDecodeOptions,
|
||||
VAEEncodeOptions,
|
||||
)
|
||||
|
||||
|
||||
def test_vae_encode_options_declares_inputs() -> None:
|
||||
"""VAE Encode (Options) exposes normal/tiled controls."""
|
||||
|
||||
inputs = VAEEncodeOptions.INPUT_TYPES()["required"]
|
||||
|
||||
assert VAEEncodeOptions.RETURN_TYPES == ("LATENT",)
|
||||
assert VAEEncodeOptions.RETURN_NAMES == ("latent",)
|
||||
assert list(inputs) == [
|
||||
"use_tiling",
|
||||
"pixels",
|
||||
"vae",
|
||||
"tile_size",
|
||||
"overlap",
|
||||
"temporal_size",
|
||||
"temporal_overlap",
|
||||
]
|
||||
assert inputs["use_tiling"][0] == "BOOLEAN"
|
||||
assert inputs["use_tiling"][1]["default"] is False
|
||||
assert inputs["pixels"][0] == "IMAGE"
|
||||
assert inputs["pixels"][1]["rawLink"] is True
|
||||
assert inputs["vae"][0] == "VAE"
|
||||
assert inputs["vae"][1]["rawLink"] is True
|
||||
_assert_tile_controls(inputs, tile_size_step=VAE_ENCODE_TILE_SIZE_STEP)
|
||||
|
||||
|
||||
def test_vae_decode_options_declares_inputs() -> None:
|
||||
"""VAE Decode (Options) exposes normal/tiled controls."""
|
||||
|
||||
inputs = VAEDecodeOptions.INPUT_TYPES()["required"]
|
||||
|
||||
assert VAEDecodeOptions.RETURN_TYPES == ("IMAGE",)
|
||||
assert VAEDecodeOptions.RETURN_NAMES == ("image",)
|
||||
assert list(inputs) == [
|
||||
"use_tiling",
|
||||
"samples",
|
||||
"vae",
|
||||
"tile_size",
|
||||
"overlap",
|
||||
"temporal_size",
|
||||
"temporal_overlap",
|
||||
]
|
||||
assert inputs["use_tiling"][0] == "BOOLEAN"
|
||||
assert inputs["use_tiling"][1]["default"] is False
|
||||
assert inputs["samples"][0] == "LATENT"
|
||||
assert inputs["samples"][1]["rawLink"] is True
|
||||
assert inputs["vae"][0] == "VAE"
|
||||
assert inputs["vae"][1]["rawLink"] is True
|
||||
_assert_tile_controls(inputs, tile_size_step=VAE_DECODE_TILE_SIZE_STEP)
|
||||
|
||||
|
||||
def test_vae_encode_options_expands_to_native_encode() -> None:
|
||||
"""Normal encode mode expands to ComfyUI's native VAEEncode."""
|
||||
|
||||
_set_graph_prefix("ENCODE_NORMAL")
|
||||
result = VAEEncodeOptions().encode(
|
||||
use_tiling=False,
|
||||
pixels=("image", 0),
|
||||
vae=("loader", 2),
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=64,
|
||||
temporal_overlap=8,
|
||||
)
|
||||
|
||||
node = _single_node(result["expand"])
|
||||
assert node["class_type"] == "VAEEncode"
|
||||
assert node["inputs"] == {"pixels": ["image", 0], "vae": ["loader", 2]}
|
||||
assert result["result"][0][0] in result["expand"]
|
||||
assert result["result"][0][1] == 0
|
||||
|
||||
|
||||
def test_vae_encode_options_expands_to_native_tiled_encode() -> None:
|
||||
"""Tiled encode mode expands to ComfyUI's native VAEEncodeTiled."""
|
||||
|
||||
_set_graph_prefix("ENCODE_TILED")
|
||||
result = VAEEncodeOptions().encode(
|
||||
use_tiling=True,
|
||||
pixels=["image", 0],
|
||||
vae=["loader", 2],
|
||||
tile_size=768,
|
||||
overlap=96,
|
||||
temporal_size=48,
|
||||
temporal_overlap=12,
|
||||
)
|
||||
|
||||
node = _single_node(result["expand"])
|
||||
assert node["class_type"] == "VAEEncodeTiled"
|
||||
assert node["inputs"] == {
|
||||
"pixels": ["image", 0],
|
||||
"vae": ["loader", 2],
|
||||
"tile_size": 768,
|
||||
"overlap": 96,
|
||||
"temporal_size": 48,
|
||||
"temporal_overlap": 12,
|
||||
}
|
||||
|
||||
|
||||
def test_vae_decode_options_expands_to_native_decode() -> None:
|
||||
"""Normal decode mode expands to ComfyUI's native VAEDecode."""
|
||||
|
||||
_set_graph_prefix("DECODE_NORMAL")
|
||||
result = VAEDecodeOptions().decode(
|
||||
use_tiling=False,
|
||||
samples=("latent", 0),
|
||||
vae=("loader", 2),
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=64,
|
||||
temporal_overlap=8,
|
||||
)
|
||||
|
||||
node = _single_node(result["expand"])
|
||||
assert node["class_type"] == "VAEDecode"
|
||||
assert node["inputs"] == {"samples": ["latent", 0], "vae": ["loader", 2]}
|
||||
assert result["result"][0][0] in result["expand"]
|
||||
assert result["result"][0][1] == 0
|
||||
|
||||
|
||||
def test_vae_decode_options_expands_to_native_tiled_decode() -> None:
|
||||
"""Tiled decode mode expands to ComfyUI's native VAEDecodeTiled."""
|
||||
|
||||
_set_graph_prefix("DECODE_TILED")
|
||||
result = VAEDecodeOptions().decode(
|
||||
use_tiling=True,
|
||||
samples=["latent", 0],
|
||||
vae=["loader", 2],
|
||||
tile_size=1024,
|
||||
overlap=128,
|
||||
temporal_size=96,
|
||||
temporal_overlap=16,
|
||||
)
|
||||
|
||||
node = _single_node(result["expand"])
|
||||
assert node["class_type"] == "VAEDecodeTiled"
|
||||
assert node["inputs"] == {
|
||||
"samples": ["latent", 0],
|
||||
"vae": ["loader", 2],
|
||||
"tile_size": 1024,
|
||||
"overlap": 128,
|
||||
"temporal_size": 96,
|
||||
"temporal_overlap": 16,
|
||||
}
|
||||
|
||||
|
||||
def _assert_tile_controls(
|
||||
inputs: dict[str, tuple[Any, ...]],
|
||||
*,
|
||||
tile_size_step: int,
|
||||
) -> None:
|
||||
"""Assert shared VAE tile control metadata."""
|
||||
|
||||
tile_size = inputs["tile_size"][1]
|
||||
assert tile_size["default"] == VAE_TILE_SIZE_DEFAULT
|
||||
assert tile_size["min"] == VAE_TILE_SIZE_MIN
|
||||
assert tile_size["max"] == VAE_TILE_SIZE_MAX
|
||||
assert tile_size["step"] == tile_size_step
|
||||
assert tile_size["advanced"] is True
|
||||
|
||||
overlap = inputs["overlap"][1]
|
||||
assert overlap["default"] == VAE_OVERLAP_DEFAULT
|
||||
assert overlap["min"] == VAE_OVERLAP_MIN
|
||||
assert overlap["max"] == VAE_OVERLAP_MAX
|
||||
assert overlap["step"] == VAE_OVERLAP_STEP
|
||||
assert overlap["advanced"] is True
|
||||
|
||||
temporal_size = inputs["temporal_size"][1]
|
||||
assert temporal_size["default"] == VAE_TEMPORAL_SIZE_DEFAULT
|
||||
assert temporal_size["min"] == VAE_TEMPORAL_SIZE_MIN
|
||||
assert temporal_size["max"] == VAE_TEMPORAL_SIZE_MAX
|
||||
assert temporal_size["step"] == VAE_TEMPORAL_SIZE_STEP
|
||||
assert temporal_size["advanced"] is True
|
||||
|
||||
temporal_overlap = inputs["temporal_overlap"][1]
|
||||
assert temporal_overlap["default"] == VAE_TEMPORAL_OVERLAP_DEFAULT
|
||||
assert temporal_overlap["min"] == VAE_TEMPORAL_OVERLAP_MIN
|
||||
assert temporal_overlap["max"] == VAE_TEMPORAL_OVERLAP_MAX
|
||||
assert temporal_overlap["step"] == VAE_TEMPORAL_OVERLAP_STEP
|
||||
assert temporal_overlap["advanced"] is True
|
||||
|
||||
|
||||
def _single_node(graph: dict[str, dict[str, Any]]) -> dict[str, Any]:
|
||||
"""Return the only node from a dynamic expansion graph."""
|
||||
|
||||
assert len(graph) == 1
|
||||
return next(iter(graph.values()))
|
||||
|
||||
|
||||
def _set_graph_prefix(prefix: str) -> None:
|
||||
"""Set a deterministic Comfy graph-builder prefix for assertions."""
|
||||
|
||||
graph_utils = import_module("comfy_execution.graph_utils")
|
||||
graph_utils.GraphBuilder.set_default_prefix(prefix, 0, 0)
|
||||
Vendored
+432
-3
@@ -3,11 +3,17 @@ import { app } from "../../../scripts/app.js";
|
||||
|
||||
// web/src/api.ts
|
||||
var SETTINGS_ROUTE = "/simple-syrup/settings";
|
||||
var EXTERNAL_LLM_SETTINGS_ROUTE = "/simple-syrup/external-llm/settings";
|
||||
var EXTERNAL_LLM_API_KEY_ROUTE = "/simple-syrup/external-llm/api-key";
|
||||
var EXTERNAL_LLM_MODELS_REFRESH_ROUTE = "/simple-syrup/external-llm/models/refresh";
|
||||
async function getSettings(fetchImpl = fetch) {
|
||||
const response = await fetchImpl(SETTINGS_ROUTE);
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
`Could not load SimpleSyrup settings. Backend returned ${String(response.status)}.`
|
||||
await backendErrorMessage(
|
||||
response,
|
||||
`Could not load SimpleSyrup settings. Backend returned ${String(response.status)}.`
|
||||
)
|
||||
);
|
||||
}
|
||||
return parseSettings(await response.json());
|
||||
@@ -20,7 +26,10 @@ async function saveSettings(settings, fetchImpl = fetch) {
|
||||
});
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
`Could not save SimpleSyrup settings. Backend returned ${String(response.status)}.`
|
||||
await backendErrorMessage(
|
||||
response,
|
||||
`Could not save SimpleSyrup settings. Backend returned ${String(response.status)}.`
|
||||
)
|
||||
);
|
||||
}
|
||||
return parseSettings(await response.json());
|
||||
@@ -35,19 +44,124 @@ function parseSettings(payload) {
|
||||
show_downloadable_models: payload.show_downloadable_models
|
||||
};
|
||||
}
|
||||
async function getExternalLLMSettings(fetchImpl = fetch) {
|
||||
const response = await fetchImpl(EXTERNAL_LLM_SETTINGS_ROUTE);
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
await backendErrorMessage(
|
||||
response,
|
||||
`Could not load external LLM settings. Backend returned ${String(response.status)}.`
|
||||
)
|
||||
);
|
||||
}
|
||||
return parseExternalLLMSettings(await response.json());
|
||||
}
|
||||
async function saveExternalLLMSettings(settings, fetchImpl = fetch) {
|
||||
const response = await fetchImpl(EXTERNAL_LLM_SETTINGS_ROUTE, {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify(settings)
|
||||
});
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
await backendErrorMessage(
|
||||
response,
|
||||
`Could not save external LLM settings. Backend returned ${String(response.status)}.`
|
||||
)
|
||||
);
|
||||
}
|
||||
return parseExternalLLMSettings(await response.json());
|
||||
}
|
||||
async function saveExternalLLMApiKey(payload, fetchImpl = fetch) {
|
||||
const response = await fetchImpl(EXTERNAL_LLM_API_KEY_ROUTE, {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify(payload)
|
||||
});
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
await backendErrorMessage(
|
||||
response,
|
||||
`Could not save external LLM API key. Backend returned ${String(response.status)}.`
|
||||
)
|
||||
);
|
||||
}
|
||||
return parseExternalLLMSettings(await response.json());
|
||||
}
|
||||
async function refreshExternalLLMModels(fetchImpl = fetch) {
|
||||
const response = await fetchImpl(EXTERNAL_LLM_MODELS_REFRESH_ROUTE, {
|
||||
method: "POST"
|
||||
});
|
||||
if (!response.ok) {
|
||||
throw new Error(
|
||||
await backendErrorMessage(
|
||||
response,
|
||||
`Could not refresh external LLM models. Backend returned ${String(response.status)}.`
|
||||
)
|
||||
);
|
||||
}
|
||||
return parseExternalLLMSettings(await response.json());
|
||||
}
|
||||
function parseExternalLLMSettings(payload) {
|
||||
if (!isExternalLLMSettingsPayload(payload)) {
|
||||
throw new Error(
|
||||
"External LLM settings payload is invalid. Expected base_url, cached_models, default_model, and has_api_key."
|
||||
);
|
||||
}
|
||||
return {
|
||||
base_url: payload.base_url,
|
||||
cached_models: [...payload.cached_models],
|
||||
default_model: payload.default_model,
|
||||
has_api_key: payload.has_api_key
|
||||
};
|
||||
}
|
||||
function isSettingsPayload(payload) {
|
||||
return typeof payload === "object" && payload !== null && typeof payload.show_downloadable_models === "boolean";
|
||||
}
|
||||
function isExternalLLMSettingsPayload(payload) {
|
||||
return typeof payload === "object" && payload !== null && typeof payload.base_url === "string" && Array.isArray(payload.cached_models) && payload.cached_models?.every(
|
||||
(model) => typeof model === "string"
|
||||
) === true && typeof payload.default_model === "string" && typeof payload.has_api_key === "boolean";
|
||||
}
|
||||
async function backendErrorMessage(response, fallback) {
|
||||
try {
|
||||
const payload = await response.json();
|
||||
if (typeof payload === "object" && payload !== null && typeof payload.error === "string") {
|
||||
return payload.error;
|
||||
}
|
||||
} catch {
|
||||
return fallback;
|
||||
}
|
||||
return fallback;
|
||||
}
|
||||
|
||||
// web/src/settings.ts
|
||||
var SIMPLE_SYRUP_SETTING_ID = "SimpleSyrup.ShowDownloadableModels";
|
||||
var SIMPLE_SYRUP_SETTING_LABEL = "SimpleSyrup: Show downloadable models in loader dropdowns";
|
||||
var SIMPLE_SYRUP_SETTING_DESCRIPTION = "Show known downloadable SAM, GroundingDINO, and ViTMatte models even when they are not installed locally.";
|
||||
var EXTERNAL_LLM_ENDPOINT_SETTING_ID = "SimpleSyrup.ExternalLLM.Endpoint";
|
||||
var EXTERNAL_LLM_ENDPOINT_SETTING_LABEL = "SimpleSyrup: External LLM endpoint";
|
||||
var EXTERNAL_LLM_ENDPOINT_SETTING_DESCRIPTION = "OpenAI-compatible endpoint base URL used by SimpleSyrup prompt nodes.";
|
||||
var EXTERNAL_LLM_API_KEY_SETTING_ID = "SimpleSyrup.ExternalLLM.ApiKey";
|
||||
var EXTERNAL_LLM_API_KEY_SETTING_LABEL = "SimpleSyrup: External LLM API key";
|
||||
var EXTERNAL_LLM_API_KEY_SETTING_DESCRIPTION = "Stores the API key for the configured external LLM endpoint in OS credential storage.";
|
||||
var DEFAULT_SETTINGS = {
|
||||
show_downloadable_models: true
|
||||
};
|
||||
async function registerSimpleSyrupSettings(app2, api = { getSettings, saveSettings }, logger = console) {
|
||||
async function registerSimpleSyrupSettings(app2, api = {
|
||||
getSettings,
|
||||
saveSettings,
|
||||
getExternalLLMSettings,
|
||||
saveExternalLLMSettings,
|
||||
saveExternalLLMApiKey
|
||||
}, logger = console) {
|
||||
let initialSettings = DEFAULT_SETTINGS;
|
||||
let externalLLMSettings = {
|
||||
base_url: "",
|
||||
cached_models: [],
|
||||
default_model: "",
|
||||
has_api_key: false
|
||||
};
|
||||
try {
|
||||
initialSettings = await api.getSettings();
|
||||
} catch (error) {
|
||||
@@ -56,7 +170,17 @@ async function registerSimpleSyrupSettings(app2, api = { getSettings, saveSettin
|
||||
error
|
||||
);
|
||||
}
|
||||
try {
|
||||
externalLLMSettings = await api.getExternalLLMSettings();
|
||||
} catch (error) {
|
||||
logger.warn(
|
||||
"Could not load SimpleSyrup external LLM settings. Using empty endpoint settings until the backend is available.",
|
||||
error
|
||||
);
|
||||
}
|
||||
let savedSettings = initialSettings;
|
||||
let savedExternalLLMSettings = externalLLMSettings;
|
||||
installSimpleSyrupSettingsStyle();
|
||||
const setting = app2.ui.settings.addSetting({
|
||||
id: SIMPLE_SYRUP_SETTING_ID,
|
||||
name: SIMPLE_SYRUP_SETTING_LABEL,
|
||||
@@ -80,6 +204,310 @@ async function registerSimpleSyrupSettings(app2, api = { getSettings, saveSettin
|
||||
}
|
||||
});
|
||||
setting.value = initialSettings.show_downloadable_models;
|
||||
app2.ui.settings.addSetting({
|
||||
id: EXTERNAL_LLM_ENDPOINT_SETTING_ID,
|
||||
name: EXTERNAL_LLM_ENDPOINT_SETTING_LABEL,
|
||||
sortOrder: 320,
|
||||
type: () => createExternalLLMEndpointControl({
|
||||
api,
|
||||
logger,
|
||||
refreshModelChoices: () => refreshExternalLLMModelChoices(app2, logger),
|
||||
getSettings: () => savedExternalLLMSettings,
|
||||
setSettings: (settings) => {
|
||||
savedExternalLLMSettings = settings;
|
||||
}
|
||||
}),
|
||||
defaultValue: externalLLMSettings.base_url,
|
||||
tooltip: EXTERNAL_LLM_ENDPOINT_SETTING_DESCRIPTION
|
||||
});
|
||||
app2.ui.settings.addSetting({
|
||||
id: EXTERNAL_LLM_API_KEY_SETTING_ID,
|
||||
name: EXTERNAL_LLM_API_KEY_SETTING_LABEL,
|
||||
sortOrder: 319,
|
||||
type: () => createExternalLLMApiKeyControl({
|
||||
api,
|
||||
logger,
|
||||
refreshModelChoices: () => refreshExternalLLMModelChoices(app2, logger),
|
||||
getSettings: () => savedExternalLLMSettings,
|
||||
setSettings: (settings) => {
|
||||
savedExternalLLMSettings = settings;
|
||||
}
|
||||
}),
|
||||
defaultValue: "",
|
||||
tooltip: EXTERNAL_LLM_API_KEY_SETTING_DESCRIPTION
|
||||
});
|
||||
}
|
||||
function endpointShouldBeSaved(value) {
|
||||
const endpoint = value.trim();
|
||||
if (!endpoint) {
|
||||
return true;
|
||||
}
|
||||
try {
|
||||
const parsed = new URL(endpoint);
|
||||
return (parsed.protocol === "http:" || parsed.protocol === "https:") && parsed.hostname.length > 0;
|
||||
} catch {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
function installSimpleSyrupSettingsStyle() {
|
||||
if (document.getElementById("simple-syrup-settings-style")) {
|
||||
return;
|
||||
}
|
||||
const style = document.createElement("style");
|
||||
style.id = "simple-syrup-settings-style";
|
||||
style.textContent = `
|
||||
.simple-syrup-settings-row {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 0.5rem;
|
||||
min-width: min(38rem, 100%);
|
||||
}
|
||||
.simple-syrup-settings-row[data-pending="true"] {
|
||||
opacity: 0.75;
|
||||
}
|
||||
.simple-syrup-settings-input {
|
||||
min-width: 16rem;
|
||||
flex: 1 1 auto;
|
||||
}
|
||||
.simple-syrup-settings-button {
|
||||
flex: 0 0 auto;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.simple-syrup-settings-status {
|
||||
color: var(--fg-color);
|
||||
opacity: 0.8;
|
||||
white-space: normal;
|
||||
overflow-wrap: anywhere;
|
||||
}
|
||||
.simple-syrup-dialog-backdrop {
|
||||
position: fixed;
|
||||
inset: 0;
|
||||
z-index: 2147483647;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
background: rgb(0 0 0 / 45%);
|
||||
}
|
||||
.simple-syrup-dialog {
|
||||
display: grid;
|
||||
gap: 0.75rem;
|
||||
min-width: min(28rem, calc(100vw - 2rem));
|
||||
padding: 1rem;
|
||||
background: var(--comfy-menu-bg);
|
||||
color: var(--fg-color);
|
||||
}
|
||||
.simple-syrup-dialog-title {
|
||||
margin: 0;
|
||||
font-size: 1rem;
|
||||
}
|
||||
.simple-syrup-dialog-actions {
|
||||
display: flex;
|
||||
justify-content: flex-end;
|
||||
gap: 0.5rem;
|
||||
}
|
||||
`;
|
||||
document.head.appendChild(style);
|
||||
}
|
||||
function createExternalLLMEndpointControl(context) {
|
||||
const wrapper = createElement("div", "simple-syrup-settings-row");
|
||||
const input = createElement("input", "simple-syrup-settings-input");
|
||||
input.type = "text";
|
||||
input.value = context.getSettings().base_url;
|
||||
input.autocomplete = "off";
|
||||
const saveButton = createElement("button", "simple-syrup-settings-button");
|
||||
saveButton.type = "button";
|
||||
saveButton.textContent = "Save Endpoint";
|
||||
const status = createElement("span", "simple-syrup-settings-status");
|
||||
const saveEndpoint = async () => {
|
||||
const value = input.value.trim();
|
||||
if (!endpointShouldBeSaved(value)) {
|
||||
status.textContent = "Enter an http:// or https:// endpoint.";
|
||||
return;
|
||||
}
|
||||
setPending(wrapper, true);
|
||||
try {
|
||||
const saved = await context.api.saveExternalLLMSettings({
|
||||
base_url: value,
|
||||
default_model: context.getSettings().default_model
|
||||
});
|
||||
context.setSettings(saved);
|
||||
input.value = saved.base_url;
|
||||
await context.refreshModelChoices();
|
||||
status.textContent = "Endpoint saved.";
|
||||
} catch (error) {
|
||||
context.logger.warn(
|
||||
"Could not save SimpleSyrup external LLM endpoint settings. The backend rejected the setting update.",
|
||||
error
|
||||
);
|
||||
status.textContent = errorMessage(error, "Endpoint was not saved.");
|
||||
} finally {
|
||||
setPending(wrapper, false);
|
||||
}
|
||||
};
|
||||
saveButton.addEventListener("click", () => {
|
||||
void saveEndpoint();
|
||||
});
|
||||
input.addEventListener("keydown", (event) => {
|
||||
if (event.key === "Enter") {
|
||||
event.preventDefault();
|
||||
void saveEndpoint();
|
||||
}
|
||||
});
|
||||
wrapper.append(input, saveButton, status);
|
||||
return wrapper;
|
||||
}
|
||||
function createExternalLLMApiKeyControl(context) {
|
||||
const wrapper = createElement("div", "simple-syrup-settings-row");
|
||||
const button = createElement("button", "simple-syrup-settings-button");
|
||||
button.type = "button";
|
||||
const status = createElement("span", "simple-syrup-settings-status");
|
||||
const render = () => {
|
||||
const remembered = context.getSettings().has_api_key;
|
||||
button.textContent = remembered ? "Replace API Key" : "Add API Key";
|
||||
status.textContent = remembered ? "API key remembered." : "";
|
||||
};
|
||||
button.addEventListener("click", () => {
|
||||
if (!context.getSettings().base_url.trim()) {
|
||||
status.textContent = "Save endpoint first.";
|
||||
return;
|
||||
}
|
||||
openExternalLLMApiKeyDialog({
|
||||
replacing: context.getSettings().has_api_key,
|
||||
onSubmit: async (apiKey) => {
|
||||
setPending(wrapper, true);
|
||||
try {
|
||||
const saved = await context.api.saveExternalLLMApiKey({
|
||||
api_key: apiKey
|
||||
});
|
||||
context.setSettings(saved);
|
||||
await context.refreshModelChoices();
|
||||
render();
|
||||
status.textContent = context.getSettings().has_api_key ? "API key remembered." : "API key was not saved.";
|
||||
} catch (error) {
|
||||
context.logger.warn(
|
||||
"Could not save SimpleSyrup external LLM API key. The backend rejected the credential update.",
|
||||
error
|
||||
);
|
||||
status.textContent = apiKeyErrorMessage(error);
|
||||
} finally {
|
||||
setPending(wrapper, false);
|
||||
}
|
||||
}
|
||||
});
|
||||
});
|
||||
wrapper.append(button, status);
|
||||
render();
|
||||
return wrapper;
|
||||
}
|
||||
function openExternalLLMApiKeyDialog(options) {
|
||||
const overlay = createElement("div", "simple-syrup-dialog-backdrop");
|
||||
const dialog = createElement("div", "simple-syrup-dialog comfy-dialog");
|
||||
const title = createElement("h3", "simple-syrup-dialog-title");
|
||||
title.textContent = options.replacing ? "Replace API Key" : "Add API Key";
|
||||
const input = createElement("input", "simple-syrup-settings-input");
|
||||
input.type = "password";
|
||||
input.autocomplete = "off";
|
||||
input.placeholder = "API key";
|
||||
input.setAttribute("data-1p-ignore", "true");
|
||||
input.setAttribute("data-lpignore", "true");
|
||||
input.setAttribute("data-bwignore", "true");
|
||||
const actions = createElement("div", "simple-syrup-dialog-actions");
|
||||
const submitButton = createElement("button", "simple-syrup-settings-button");
|
||||
submitButton.type = "button";
|
||||
submitButton.textContent = options.replacing ? "Replace Key" : "Store Key";
|
||||
const cancelButton = createElement("button", "simple-syrup-settings-button");
|
||||
cancelButton.type = "button";
|
||||
cancelButton.textContent = "Cancel";
|
||||
const close = () => {
|
||||
overlay.remove();
|
||||
};
|
||||
const submit = async () => {
|
||||
const apiKey = input.value.trim();
|
||||
if (!apiKey) {
|
||||
input.focus();
|
||||
return;
|
||||
}
|
||||
submitButton.disabled = true;
|
||||
await options.onSubmit(apiKey);
|
||||
close();
|
||||
};
|
||||
submitButton.addEventListener("click", () => {
|
||||
void submit();
|
||||
});
|
||||
cancelButton.addEventListener("click", close);
|
||||
input.addEventListener("keydown", (event) => {
|
||||
if (event.key === "Enter") {
|
||||
event.preventDefault();
|
||||
void submit();
|
||||
}
|
||||
if (event.key === "Escape") {
|
||||
event.preventDefault();
|
||||
close();
|
||||
}
|
||||
});
|
||||
actions.append(submitButton, cancelButton);
|
||||
dialog.append(title, input, actions);
|
||||
overlay.append(dialog);
|
||||
document.body.appendChild(overlay);
|
||||
input.focus();
|
||||
}
|
||||
function createElement(tagName, className) {
|
||||
const element = document.createElement(tagName);
|
||||
element.className = className;
|
||||
return element;
|
||||
}
|
||||
function setPending(element, pending) {
|
||||
element.dataset.pending = pending ? "true" : "false";
|
||||
for (const control of Array.from(element.querySelectorAll("input, button"))) {
|
||||
if (control instanceof HTMLInputElement || control instanceof HTMLButtonElement) {
|
||||
control.disabled = pending;
|
||||
}
|
||||
}
|
||||
}
|
||||
function errorMessage(error, fallback) {
|
||||
if (error instanceof Error && error.message) {
|
||||
return error.message;
|
||||
}
|
||||
return fallback;
|
||||
}
|
||||
function apiKeyErrorMessage(error) {
|
||||
const message = errorMessage(error, "API key was not saved.");
|
||||
if (message.includes("Configure an external LLM endpoint")) {
|
||||
return "Save endpoint first.";
|
||||
}
|
||||
return message;
|
||||
}
|
||||
async function refreshExternalLLMModelChoices(app2, logger) {
|
||||
try {
|
||||
await app2.refreshComboInNodes?.();
|
||||
} catch (error) {
|
||||
logger.warn(
|
||||
"Could not refresh Comfy node definitions after saving external LLM settings.",
|
||||
error
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// web/src/refresh.ts
|
||||
var REFRESH_WRAPPED = /* @__PURE__ */ Symbol.for("SimpleSyrup.ExternalLLM.RefreshWrapped");
|
||||
function registerExternalLLMRefreshHook(app2, api = { refreshExternalLLMModels }, logger = console) {
|
||||
if (!app2.refreshComboInNodes) {
|
||||
return;
|
||||
}
|
||||
const refreshOwner = app2;
|
||||
if (refreshOwner[REFRESH_WRAPPED]) {
|
||||
return;
|
||||
}
|
||||
const originalRefresh = app2.refreshComboInNodes.bind(app2);
|
||||
refreshOwner[REFRESH_WRAPPED] = true;
|
||||
app2.refreshComboInNodes = async () => {
|
||||
try {
|
||||
await api.refreshExternalLLMModels();
|
||||
} catch (error) {
|
||||
logger.warn("Could not refresh SimpleSyrup external LLM models.", error);
|
||||
}
|
||||
await originalRefresh();
|
||||
};
|
||||
}
|
||||
|
||||
// web/src/main.ts
|
||||
@@ -88,5 +516,6 @@ comfyApp.registerExtension({
|
||||
name: "SimpleSyrup.Settings",
|
||||
async setup(appInstance) {
|
||||
await registerSimpleSyrupSettings(appInstance);
|
||||
registerExternalLLMRefreshHook(appInstance);
|
||||
}
|
||||
});
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user