diff --git a/.github/workflows/quality-gates.yml b/.github/workflows/quality-gates.yml new file mode 100644 index 0000000..546c804 --- /dev/null +++ b/.github/workflows/quality-gates.yml @@ -0,0 +1,62 @@ +name: quality gates + +on: + pull_request: + +permissions: + contents: read + +jobs: + quality: + runs-on: ubuntu-latest + env: + SIMPLE_SYRUP_TEST_COMFY_CPU: "1" + steps: + - name: Checkout + uses: actions/checkout@v6 + + - name: Setup Node + uses: actions/setup-node@v6 + with: + node-version: 22.14.0 + cache: npm + + - name: Setup Python + uses: actions/setup-python@v6 + with: + python-version: "3.11" + + - name: Install Node dependencies + run: npm ci + + - name: Install ComfyUI host dependencies + shell: bash + run: | + set -euo pipefail + git clone --depth 1 https://github.com/comfyanonymous/ComfyUI.git "$RUNNER_TEMP/ComfyUI" + rsync -a --exclude=".git" "$RUNNER_TEMP/ComfyUI/" "$GITHUB_WORKSPACE/../.."/ + pip install -r "$GITHUB_WORKSPACE/../../requirements.txt" + + - name: Install Python dependencies + run: pip install -e . pytest pytest-xdist ruff mypy + + - name: Check architecture governance + run: python -m tools.check_architecture + + - name: Check test governance + run: python -m tools.check_test_governance + + - name: Verify Python formatting + run: ruff format --check . + + - name: Verify Python lint + run: ruff check . + + - name: Verify Python types + run: mypy --strict simple_syrup tests + + - name: Verify Python tests + run: pytest -n auto -q -m "not external_artifact" + + - name: Verify frontend + run: npm run check:web diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index a684f7e..5b8f406 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -55,8 +55,8 @@ jobs: --noconftest --rootdir=tests --confcutdir=tests - tests/test_grounding_dino_bert_adapter.py - tests/test_grounding_dino_text_token_masks.py + tests/segmentation/detection/test_grounding_dino_bert_adapter.py + tests/segmentation/detection/test_grounding_dino_text_token_masks.py release: if: github.event_name != 'pull_request' @@ -98,6 +98,12 @@ jobs: - name: Verify Python formatting run: ruff format --check . + - name: Check architecture governance + run: python -m tools.check_architecture + + - name: Check test governance + run: python -m tools.check_test_governance + - name: Verify Python lint run: ruff check . diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml new file mode 100644 index 0000000..c271f1f --- /dev/null +++ b/.pre-commit-config.yaml @@ -0,0 +1,29 @@ +repos: + - repo: local + hooks: + - id: architecture-governance + name: Enforce architecture governance + entry: ..\..\venv\Scripts\python.exe -m tools.check_architecture + language: system + pass_filenames: false + always_run: true + - id: test-governance + name: Enforce test governance + entry: ..\..\venv\Scripts\python.exe -m tools.check_test_governance + language: system + pass_filenames: false + always_run: true + + - repo: https://github.com/pre-commit/pre-commit-hooks + rev: v5.0.0 + hooks: + - id: end-of-file-fixer + - id: mixed-line-ending + args: [--fix=lf] + exclude: '(\.bat$|\.cmd$|\.ps1$)' + - id: trailing-whitespace + - id: check-merge-conflict + - id: check-yaml + - id: check-json + exclude: '(^web/dist/|tsconfig\.json$)' + - id: check-toml diff --git a/AGENTS.md b/AGENTS.md index e2807a8..8694d59 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -35,6 +35,8 @@ Engineering priority is strict architecture, strong separation of concerns, comp ### Required Command Forms - Tests: `..\..\venv\Scripts\python.exe -m pytest -n auto -q` +- Architecture: `..\..\venv\Scripts\python.exe -m tools.check_architecture` +- Test governance: `..\..\venv\Scripts\python.exe -m tools.check_test_governance` - Lint: `..\..\venv\Scripts\ruff.exe check .` - Format: `..\..\venv\Scripts\ruff.exe format .` - Type check: `..\..\venv\Scripts\mypy.exe --strict simple_syrup tests` @@ -94,6 +96,34 @@ If a required tool is missing from `..\..\venv`, install or update development d - Reorganize modules when it improves architecture. - Align touched modules with the ownership and dependency rules in this file. +## Architecture Governance + +- Repository governance lives under `governance/`. +- `governance/architecture/policy.toml` defines every authored-code root, + extension, exclusion, and the 350-line soft and 500-line hard structural + thresholds. +- `governance/architecture/debt.toml` records exact assessed mixed ownership. +- `governance/architecture/waivers.toml` records exact bounded hard-gate + exceptions. +- `governance/architecture/import_debt.toml` records exact current dependency- + direction violations; new violations are prohibited. +- `governance/architecture/soft_reviews.toml` records the current human + disposition of every file between the soft and hard thresholds. +- Every hard-gate file requires source-level ownership review. +- Use a structural waiver only for one cohesive authoritative owner whose + invariants would be divided by extraction. +- Mixed ownership requires debt and a linked remediation waiver naming the + next extraction and a lower next limit. +- Waivers and debt are fingerprinted current state, not historical ledgers. +- Delete resolved records; do not extend dates or limits merely to pass the + checker. +- `governance/testing/policy.toml` defines Python and frontend test-layout and + reliability discovery. +- Every test-governance candidate requires an exact classification or + debt-remediation disposition. +- Run both governance checkers after changing authored structure, test + placement, isolation, timing, resources, or reviewed state. + ## ComfyUI Node Rules - Public node identifiers are compatibility-sensitive. diff --git a/__init__.py b/__init__.py index 69a946e..5870185 100644 --- a/__init__.py +++ b/__init__.py @@ -12,22 +12,24 @@ from . import simple_syrup as _simple_syrup_package sys.modules.setdefault("simple_syrup", _simple_syrup_package) +from .simple_syrup.integration.external_llm_routes import ( # noqa: E402 + register_external_llm_routes, +) +from .simple_syrup.integration.mask_batch_preview_routes import ( # noqa: E402 + register_mask_batch_preview_routes, +) +from .simple_syrup.integration.quant_cache_routes import ( # noqa: E402 + register_quant_cache_routes, +) +from .simple_syrup.integration.settings_routes import ( # noqa: E402 + register_settings_routes, +) from .simple_syrup.runtime.attention_region_prompt_handler import ( # noqa: E402 register_attention_region_prompt_handler, ) from .simple_syrup.runtime.comfy_safetensors_dtypes import ( # noqa: E402 register_comfy_safetensors_dtypes, ) -from .simple_syrup.runtime.external_llm_routes import ( # noqa: E402 - register_external_llm_routes, -) -from .simple_syrup.runtime.mask_batch_preview_routes import ( # noqa: E402 - register_mask_batch_preview_routes, -) -from .simple_syrup.runtime.quant_cache_routes import ( # noqa: E402 - register_quant_cache_routes, -) -from .simple_syrup.runtime.settings_routes import register_settings_routes # noqa: E402 WEB_DIRECTORY = "./web/dist" diff --git a/governance/architecture/debt.toml b/governance/architecture/debt.toml new file mode 100644 index 0000000..1297fe7 --- /dev/null +++ b/governance/architecture/debt.toml @@ -0,0 +1,2 @@ +schema_version = 1 +debts = [] diff --git a/governance/architecture/import_debt.toml b/governance/architecture/import_debt.toml new file mode 100644 index 0000000..a988bf3 --- /dev/null +++ b/governance/architecture/import_debt.toml @@ -0,0 +1 @@ +schema_version = 1 diff --git a/governance/architecture/policy.toml b/governance/architecture/policy.toml new file mode 100644 index 0000000..aa15e5e --- /dev/null +++ b/governance/architecture/policy.toml @@ -0,0 +1,41 @@ +schema_version = 2 + +[structure] +soft_lines = 350 +hard_lines = 500 +source_roots = [ + "simple_syrup/domain", + "simple_syrup/image", + "simple_syrup/integration", + "simple_syrup/masking", + "simple_syrup/nodes", + "simple_syrup/nodes_v3", + "simple_syrup/runtime", + "simple_syrup/services", + "simple_syrup/shared", + "tests", + "tools", + "scripts", + "web/src", + "web/tests", +] +source_files = [ + "__init__.py", + ".releaserc.cjs", + "eslint.config.js", + "simple_syrup/__init__.py", + "vitest.config.ts", +] +source_extensions = [ + ".cjs", + ".js", + ".mjs", + ".py", + ".pyi", + ".ts", +] +excluded_paths = [] + +[registries] +debt = "governance/architecture/debt.toml" +waivers = "governance/architecture/waivers.toml" diff --git a/governance/architecture/soft_reviews.toml b/governance/architecture/soft_reviews.toml new file mode 100644 index 0000000..45c7686 --- /dev/null +++ b/governance/architecture/soft_reviews.toml @@ -0,0 +1,59 @@ +schema_version = 1 +review_by = 2027-03-31 +fingerprint = "sha256:0a365aa7572ed38582a6fb4f09acefe700212ba65860e93e8c7cd8d5d7d12c45" + +cohesive_paths = [ + "simple_syrup/masking/prompt_segs_with_sam_service.py", + "simple_syrup/nodes/prompt_segs_with_sam.py", + "simple_syrup/nodes_v3/legacy_node_wrappers.py", + "simple_syrup/runtime/attention_region_affinity.py", + "simple_syrup/runtime/attention_region_capture.py", + "simple_syrup/runtime/attention_sampler_lineage.py", + "simple_syrup/runtime/regional_lora/anima_module_surface.py", + "simple_syrup/runtime/spatial_model_arguments.py", + "simple_syrup/services/concept_attention_evidence.py", + "tests/comfy_integration/test_comfy_regional_adapter_resolver.py", + "tests/comfy_integration/test_comfy_regional_conditioning_processing.py", + "tests/models/loading/test_checkpoint_quantizer.py", + "tests/models/patching/test_model_patcher_mutations.py", + "tests/prompting/prompt_control/test_prompt_control_schedule_encode_graph.py", + "tests/regional_generation/anima/test_anima_activation_context.py", + "tests/regional_generation/anima/test_anima_full_tile_lora_equivalence.py", + "tests/regional_generation/anima/test_anima_loader.py", + "tests/regional_generation/anima/test_anima_multi_lora_composition.py", + "tests/regional_generation/anima/test_anima_regional_diagnostics.py", + "tests/regional_generation/anima/test_anima_regional_diagnostics_wrapper.py", + "tests/regional_generation/anima/test_anima_regional_permutation_diagnostics.py", + "tests/regional_generation/anima/test_anima_single_adapter_mutations.py", + "tests/regional_generation/attention_coupling/test_attention_coupling_model_preparation_service.py", + "tests/regional_generation/attention_regions/test_attention_region_capture.py", + "tests/regional_generation/attention_regions/test_attention_region_completion.py", + "tests/regional_generation/attention_regions/test_attention_region_components.py", + "tests/regional_generation/attention_regions/test_attention_region_geometry.py", + "tests/regional_generation/regional/test_regional_attention_batching.py", + "tests/regional_generation/regional/test_regional_linear_execution.py", + "tests/regional_generation/regional/test_regional_model_patch_interop.py", + "tests/regional_generation/regional/test_regional_multidiffusion_sampling.py", + "tests/regional_generation/spatial/test_contextual_model_wrapper.py", + "tests/sampling/test_multidiffusion_sampling.py", + "tests/sampling/test_sampling_scheduler_references.py", + "tests/sampling/test_sampling_schedulers.py", + "tests/segmentation/detection/test_ultralytics_loader.py", + "tests/segmentation/segs/test_detail_segs_as_regions_service.py", + "tests/segmentation/segs/test_prompt_segs_with_sam_node.py", + "tools/architecture_governance/validation.py", + "tools/attention_coupling_benchmark/comfy_probe/negpip_runtime.py", + "tools/negpip_integration/run.py", + "tools/prompt_control_attention_coupling_integration/validation.py", + "tools/run_global_prompt_lora_proof.py", + "tools/test_governance/semantic_patterns.py", + "tools/test_governance/validation.py", + "web/src/orderedMediaNode.ts", + "web/src/orderedMediaPreviewActions.ts", + "web/tests/media/orderedMediaPreviewActions.test.ts", +] + +debt_paths = [ +] + +remediations = [] diff --git a/governance/architecture/waivers.toml b/governance/architecture/waivers.toml new file mode 100644 index 0000000..8d06866 --- /dev/null +++ b/governance/architecture/waivers.toml @@ -0,0 +1,78 @@ +schema_version = 1 + +[[waivers]] +id = "SSY-WAIVER-S001" +owner = "model catalog" +rule = "STRUCT003" +path = "simple_syrup/runtime/model_catalog.py" +kind = "structural" +justification = "This module is the single immutable catalog authority for supported model families and artifacts. Most of its size is declarative checksums, repository identities, filenames, and URLs; its small query surface and entry constructors enforce one catalog schema and change with that same metadata contract. Splitting entries by provider would scatter uniqueness and lookup review without separating behavior or ownership." +issue = "chore:SSY-WAIVER-S001" +review_by = 2027-03-31 +max_lines = 687 + +[[waivers]] +id = "SSY-WAIVER-S002" +owner = "sampling scheduler policy" +rule = "STRUCT003" +path = "simple_syrup/runtime/sampling_schedulers.py" +kind = "structural" +justification = "This module owns the complete sigma-schedule policy exposed to every sampler: supported names, fixed published AYS/GITS tables, Comfy delegation, local schedule calculation, denoise truncation, and sampler-specific terminal handling. Roughly half the file is immutable numeric reference data, while the executable functions share one public calculation boundary and dependency direction." +issue = "chore:SSY-WAIVER-S002" +review_by = 2027-03-31 +max_lines = 588 + +[[waivers]] +id = "SSY-WAIVER-S003" +owner = "Anima cross-attention patch contracts" +rule = "STRUCT003" +path = "tests/regional_generation/anima/test_anima_cross_attention.py" +kind = "structural" +justification = "This module is one integration contract for AnimaRegionalCrossAttentionPatch: it installs the exact Anima module surface, supplies one deterministic attention double, drives branch/mask/context alignment, verifies failure restoration, and proves all 28 clone-local patches. The sizable builders encode a single valid execution context and are not independent production responsibilities." +issue = "chore:SSY-WAIVER-S003" +review_by = 2027-03-31 +max_lines = 661 + +[[waivers]] +id = "SSY-WAIVER-S004" +owner = "Anima multi-LoRA fidelity contracts" +rule = "STRUCT003" +path = "tests/regional_generation/anima/test_anima_multi_lora_fidelity.py" +kind = "structural" +justification = "This module owns one numerical fidelity matrix for ordered multi-LoRA composition across schedules, branches, regions, target families, and the complete Anima surface. Its execution and reference helpers intentionally remain adjacent so every permutation is compared through the same independently calculated oracle; splitting by scenario would duplicate or conceal that shared proof authority." +issue = "chore:SSY-WAIVER-S004" +review_by = 2027-03-31 +max_lines = 678 + +[[waivers]] +id = "SSY-WAIVER-S005" +owner = "regional convolution execution contracts" +rule = "STRUCT003" +path = "tests/regional_generation/regional/test_regional_convolution_execution.py" +kind = "structural" +justification = "This module is the complete numerical contract for RegionalConvolutionExecutor across direct, pointwise, LoCon, strided, grouped, tiled-batch, ordered-adapter, and low-precision execution. Its fixture builds the same execution plan and independent convolution reference for every case, so the tests share one owner, oracle, dependency surface, and change cadence." +issue = "chore:SSY-WAIVER-S005" +review_by = 2027-03-31 +max_lines = 556 + +[[waivers]] +id = "SSY-WAIVER-S006" +owner = "Ultralytics detection node contracts" +rule = "STRUCT003" +path = "tests/segmentation/detection/test_detect_segs_with_ultralytics_node.py" +kind = "structural" +justification = "This module owns the workflow-facing contract of one Comfy node, including its schema, exact input order, batch behavior, sorting/ranking limits, union mode, and output shape. The service and builder doubles are deliberately local representations of that node boundary; every test changes with the same node API and persisted workflow contract." +issue = "chore:SSY-WAIVER-S006" +review_by = 2027-03-31 +max_lines = 592 + +[[waivers]] +id = "SSY-WAIVER-S007" +owner = "scale-factor detail service contracts" +rule = "STRUCT003" +path = "tests/segmentation/segs/test_detail_segs_by_scale_factor_service.py" +kind = "structural" +justification = "This module is the end-to-end behavioral contract for DetailSEGSByScaleFactorService, whose single orchestration transaction selects per-segment conditioning, sizes and resizes crops, applies masks, samples, decodes, and pastes results. Its sampler and resizer doubles record that one transaction; splitting them would duplicate setup without creating a distinct behavior owner." +issue = "chore:SSY-WAIVER-S007" +review_by = 2027-03-31 +max_lines = 532 diff --git a/governance/testing/debt.toml b/governance/testing/debt.toml new file mode 100644 index 0000000..1297fe7 --- /dev/null +++ b/governance/testing/debt.toml @@ -0,0 +1,2 @@ +schema_version = 1 +debts = [] diff --git a/governance/testing/policy.toml b/governance/testing/policy.toml new file mode 100644 index 0000000..1cd0583 --- /dev/null +++ b/governance/testing/policy.toml @@ -0,0 +1,27 @@ +schema_version = 1 + +[scope] +test_root = "tests" +semantic_support_roots = ["tools"] +root_source_extensions = [".py", ".pyi"] +allowed_root_source_paths = [ + "tests/ci_test_policy.py", + "tests/conftest.py", +] + +[discovery] +serial_policy = "tests/ci_test_policy.py" +wait_calls = ["QTest.qWait", "time.sleep"] +wall_clock_calls = [ + "QElapsedTimer", + "monotonic", + "perf_counter", + "time.monotonic", + "time.perf_counter", +] +xdist_environment_name = "PYTEST_XDIST_WORKER" +repository_scratch_name = ".pytest-tmp" + +[registries] +debt = "governance/testing/debt.toml" +waivers = "governance/testing/waivers.toml" diff --git a/governance/testing/waivers.toml b/governance/testing/waivers.toml new file mode 100644 index 0000000..6e3841e --- /dev/null +++ b/governance/testing/waivers.toml @@ -0,0 +1,153 @@ +schema_version = 1 + +[[waivers]] +id = "SSY-TEST-WAIVER-C001" +owner = "pytest CUDA isolation bootstrap" +kind = "classification" +disposition = "framework_infrastructure" +rule = "ENV001" +candidates = ["ENV001|tests/conftest.py|:environment-mutation:1"] +paths = ["tests/conftest.py"] +fingerprint = "sha256:481069f16237eff312f817f1e4dd3213e804eee0f3341e3f6c4aa2ba73066af3" +rationale = "The root pytest bootstrap disables CUDA visibility before Torch and ComfyUI are imported unless the maintainer explicitly enables hardware tests. Every xdist worker receives the same inherited setting before collection, so this is suite framework configuration rather than mutable test-owned state." +issue = "chore:SSY-TEST-WAIVER-C001" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C002" +owner = "native checkpoint quantization proofs" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = [ + "OPTIONAL001|tests/models/loading/test_checkpoint_quantizer.py|:optional-proof:1", + "OPTIONAL001|tests/models/loading/test_checkpoint_quantizer.py|:optional-proof:2", + "OPTIONAL001|tests/models/loading/test_checkpoint_quantizer.py|:optional-proof:3", + "OPTIONAL001|tests/models/loading/test_checkpoint_quantizer.py|:optional-proof:4", + "OPTIONAL001|tests/models/loading/test_checkpoint_quantizer.py|:optional-proof:5", +] +paths = ["tests/models/loading/test_checkpoint_quantizer.py"] +fingerprint = "sha256:bc2bee06ddb6cf41ba66497855e7349e5aa5dd5292f732284a812a615fa65c49" +rationale = "These proofs exercise installed ComfyUI NVFP4/MXFP8 kernels, GPU compute capability, and optional comfy-aimdo reload behavior. CPU fake-boundary tests in the same module always run; only the native serialization contracts are skipped when their external runtime or hardware capability does not exist." +issue = "chore:SSY-TEST-WAIVER-C002" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C003" +owner = "CUDA Anima projection precision proof" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = ["OPTIONAL001|tests/regional_generation/anima/test_anima_projection_batch.py|:optional-proof:1"] +paths = ["tests/regional_generation/anima/test_anima_projection_batch.py"] +fingerprint = "sha256:cb5e280524450a58144b350799129b3c27a6d4ce430867de582da19303987c33" +rationale = "This exact comparison bounds BF16 batched projection error on CUDA tensors against independently generated projections. Its behavior depends on the installed CUDA execution path and cannot truthfully be substituted by CPU arithmetic; all device-independent projection contracts remain mandatory." +issue = "chore:SSY-TEST-WAIVER-C003" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C004" +owner = "native Anima quantization workflow" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = ["OPTIONAL001|tests/regional_generation/anima/test_anima_quantization_workflow.py|:optional-proof:1"] +paths = ["tests/regional_generation/anima/test_anima_quantization_workflow.py"] +fingerprint = "sha256:ff02416ea8c16112605b221975425a6f88013f66e2a8bc66bdc2d1a0ecfa5053" +rationale = "The workflow proof intentionally uses the installed ComfyUI NVFP4 implementation and the active GPU's native compute support before loading the generated Anima artifact. It remains optional only where that hardware capability is absent; the portable resolver and policy tests still run everywhere." +issue = "chore:SSY-TEST-WAIVER-C004" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C005" +owner = "installed Anima CUDA smoke" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = ["OPTIONAL001|tests/regional_generation/anima/test_anima_regional_model_smoke.py|:optional-proof:1"] +paths = ["tests/regional_generation/anima/test_anima_regional_model_smoke.py"] +fingerprint = "sha256:73b3bde9d6009b71608aec73dde862d963a3a723d5ad5d57c130428a4c1511c6" +rationale = "This smoke test constructs the installed Comfy Anima model and executes its complete patched forward on CUDA tensors. It proves the native device/runtime integration and is skipped only without CUDA; deterministic component and surface contracts cover the same code boundaries on every host." +issue = "chore:SSY-TEST-WAIVER-C005" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C006" +owner = "regional convolution CUDA precision" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = ["OPTIONAL001|tests/regional_generation/regional/test_regional_convolution_execution.py|:optional-proof:1"] +paths = ["tests/regional_generation/regional/test_regional_convolution_execution.py"] +fingerprint = "sha256:f197de881034f1c11b46ce290f3b6c515fe251d5bdef7efb419f43ad7257bf40" +rationale = "The optional parameterized cases prove FP16 and BF16 regional convolution behavior through the installed CUDA kernels. CPU tests in the same contract cover dimensions, grouping, stride, masking, ordering, and reference math; only device-specific low-precision execution requires CUDA." +issue = "chore:SSY-TEST-WAIVER-C006" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C007" +owner = "regional linear CUDA precision" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = [ + "OPTIONAL001|tests/regional_generation/regional/test_regional_linear_execution.py|:optional-proof:1", + "OPTIONAL001|tests/regional_generation/regional/test_regional_linear_execution.py|:optional-proof:2", +] +paths = ["tests/regional_generation/regional/test_regional_linear_execution.py"] +fingerprint = "sha256:d8a833ef160838b80db21d7240d789879deb8a4dc39f89d52145b1ebf580f765" +rationale = "These cases validate installed CUDA FP16/BF16 projection rounding and compatible-adapter accumulation on the actual device execution path. The module's CPU contracts always prove masking, ordering, preparation, and reference deltas; the classified cases add hardware-specific numerical evidence." +issue = "chore:SSY-TEST-WAIVER-C007" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C008" +owner = "fused regional LoRA CUDA kernel" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = ["OPTIONAL001|tests/regional_generation/regional/test_regional_lora_fused_active_accumulation.py|:optional-proof:1"] +paths = ["tests/regional_generation/regional/test_regional_lora_fused_active_accumulation.py"] +fingerprint = "sha256:2053ff281c4e707c99db6424074ef3526c901ad84f0f750e044c35116981cd9c" +rationale = "The entire module qualifies the CUDA-only fused regional LoRA accumulator across low-precision dtypes, adapter counts, and indexed paths. There is no CPU implementation to exercise, while the non-fused accumulation owner has mandatory portable reference coverage." +issue = "chore:SSY-TEST-WAIVER-C008" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C009" +owner = "fused multiplier CUDA transport" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = ["OPTIONAL001|tests/regional_generation/regional/test_regional_lora_fused_multiplier_transport.py|:optional-proof:1"] +paths = ["tests/regional_generation/regional/test_regional_lora_fused_multiplier_transport.py"] +fingerprint = "sha256:56ef4217551bdccc97fc6277bd2ade11ac10b9512ee3725c46b210044f380f21" +rationale = "This module proves multiple adapter multipliers reach the CUDA fused kernel without an intermediate stack. The production behavior exists only for a CUDA-capable device, and portable composition tests independently cover ordering and multiplier semantics outside this native optimization." +issue = "chore:SSY-TEST-WAIVER-C009" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C010" +owner = "ordered tensor CUDA accumulation" +kind = "classification" +disposition = "platform_native" +rule = "OPTIONAL001" +candidates = ["OPTIONAL001|tests/sampling/test_ordered_tensor_accumulation.py|:optional-proof:1"] +paths = ["tests/sampling/test_ordered_tensor_accumulation.py"] +fingerprint = "sha256:b6306323fd555706f0b7525078acfa199c657f47f96432ce320da9057fc41709" +rationale = "The parameter matrix compares stepwise CUDA accumulation and its exact low-precision rounding across one and multiple Triton launches. Mandatory CPU contracts prove ordered in-place accumulation; only the GPU kernel and device dtypes are capability-gated." +issue = "chore:SSY-TEST-WAIVER-C010" +review_by = 2027-03-31 + +[[waivers]] +id = "SSY-TEST-WAIVER-C011" +owner = "managed Windows Comfy process lifetime" +kind = "classification" +disposition = "platform_native" +rule = "PROCESS001" +candidates = ["PROCESS001|tools/comfy_integration/server_process.py|:unscoped-child-process:1"] +paths = ["tools/comfy_integration/server_process.py"] +fingerprint = "sha256:6892815fbffe96d59bbb3f9069e44bfcdb9d6ea39e3f17f844695408af4233c0" +rationale = "WindowsComfyProcess intentionally transfers the created Popen and log handles into an explicit long-lived owner because the integration run must use the server after start returns. Its stop method signals the exact process group, bounds both graceful and forced waits, retries bounded taskkill calls, and closes both logs in finally." +issue = "chore:SSY-TEST-WAIVER-C011" +review_by = 2027-03-31 diff --git a/pyproject.toml b/pyproject.toml index e16af77..c8b740a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -69,6 +69,7 @@ ignore_missing_imports = true [tool.pytest.ini_options] pythonpath = [".", "../.."] testpaths = ["tests"] +addopts = ["--strict-markers", "--import-mode=importlib"] markers = [ "external_artifact: requires a locally installed external source or generated benchmark artifact", ] diff --git a/simple_syrup/domain/attention_coupling_preparation.py b/simple_syrup/domain/attention_coupling_preparation.py new file mode 100644 index 0000000..9e4bd50 --- /dev/null +++ b/simple_syrup/domain/attention_coupling_preparation.py @@ -0,0 +1,20 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define validated Attention Coupling preparation data.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from .raw_regional_attention import RawRegionalAttentionPlan + + +@dataclass(frozen=True, slots=True) +class AttentionCouplingPreparation: + """Retain the full raw plan and base-only ordinary sampler inputs.""" + + plan: RawRegionalAttentionPlan + positive: object + negative: object diff --git a/simple_syrup/domain/seg_preview.py b/simple_syrup/domain/seg_preview.py new file mode 100644 index 0000000..be3396e --- /dev/null +++ b/simple_syrup/domain/seg_preview.py @@ -0,0 +1,51 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define renderer-neutral SEG preview documents.""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +from .segs import CropRegion + + +@dataclass(frozen=True) +class AtlasPlacement: + """Locate one region mask inside the packed mask atlas.""" + + left: int + top: int + width: int + height: int + + +@dataclass(frozen=True) +class SegPreviewRegion: + """Describe one interactive region and its packed mask geometry.""" + + region_id: str + index: int + label: str + confidence: float + active_area: int + color: str + crop: CropRegion + atlas: AtlasPlacement + + +@dataclass(frozen=True) +class SegPreviewDocument: + """Carry bounded image assets and interaction metadata to a UI adapter.""" + + source_width: int + source_height: int + preview_width: int + preview_height: int + image: torch.Tensor + atlas: torch.Tensor + region_images: tuple[torch.Tensor, ...] + regions: tuple[SegPreviewRegion, ...] diff --git a/simple_syrup/masking/segs_mask_ops.py b/simple_syrup/domain/segs_mask_ops.py similarity index 99% rename from simple_syrup/masking/segs_mask_ops.py rename to simple_syrup/domain/segs_mask_ops.py index 07253ea..dc25422 100644 --- a/simple_syrup/masking/segs_mask_ops.py +++ b/simple_syrup/domain/segs_mask_ops.py @@ -9,7 +9,7 @@ from __future__ import annotations import torch import torch.nn.functional as F -from ..domain.segs import BoundingBox, CropRegion +from .segs import BoundingBox, CropRegion def validate_single_image(image: object, operation: str) -> torch.Tensor: diff --git a/simple_syrup/domain/tile_segs.py b/simple_syrup/domain/tile_segs.py index 6b340f6..9e1b736 100644 --- a/simple_syrup/domain/tile_segs.py +++ b/simple_syrup/domain/tile_segs.py @@ -11,11 +11,6 @@ from dataclasses import dataclass import torch -from ..masking.segs_mask_ops import ( - crop_region_for_bbox, - resize_mask, - validate_single_image, -) from ..shared.logging import get_logger from .segs import ( BoundingBox, @@ -23,6 +18,11 @@ from .segs import ( NativeSegs, Segment, ) +from .segs_mask_ops import ( + crop_region_for_bbox, + resize_mask, + validate_single_image, +) LOGGER = get_logger(__name__) diff --git a/simple_syrup/integration/__init__.py b/simple_syrup/integration/__init__.py new file mode 100644 index 0000000..d074ea4 --- /dev/null +++ b/simple_syrup/integration/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own inbound ComfyUI integration and transport composition.""" diff --git a/simple_syrup/runtime/external_llm_routes.py b/simple_syrup/integration/external_llm_routes.py similarity index 99% rename from simple_syrup/runtime/external_llm_routes.py rename to simple_syrup/integration/external_llm_routes.py index 40a0dfe..426dba1 100644 --- a/simple_syrup/runtime/external_llm_routes.py +++ b/simple_syrup/integration/external_llm_routes.py @@ -14,9 +14,9 @@ from typing import Any, Protocol, cast from aiohttp import web from ..domain.external_llm import ExternalLLMConfigError, ExternalLLMProviderError +from ..runtime.external_llm_keyring import ExternalLLMKeyringError 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" diff --git a/simple_syrup/runtime/mask_batch_preview_routes.py b/simple_syrup/integration/mask_batch_preview_routes.py similarity index 100% rename from simple_syrup/runtime/mask_batch_preview_routes.py rename to simple_syrup/integration/mask_batch_preview_routes.py diff --git a/simple_syrup/runtime/quant_cache_routes.py b/simple_syrup/integration/quant_cache_routes.py similarity index 98% rename from simple_syrup/runtime/quant_cache_routes.py rename to simple_syrup/integration/quant_cache_routes.py index f48eaae..ce1f5df 100644 --- a/simple_syrup/runtime/quant_cache_routes.py +++ b/simple_syrup/integration/quant_cache_routes.py @@ -12,13 +12,13 @@ from typing import Any, Protocol, cast from aiohttp import web +from ..runtime.quant_cache_settings import SettingsQuantCacheLimitProvider from ..services.quant_cache_service import ( QuantCacheEvictionResult, QuantCacheService, QuantCacheStatus, ) from ..services.quantized_model_boundaries import QuantCacheLimitProvider -from .quant_cache_settings import SettingsQuantCacheLimitProvider QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache" Handler = Callable[[Any], Coroutine[Any, Any, web.Response]] diff --git a/simple_syrup/runtime/settings_routes.py b/simple_syrup/integration/settings_routes.py similarity index 97% rename from simple_syrup/runtime/settings_routes.py rename to simple_syrup/integration/settings_routes.py index d6f1b23..a4eb878 100644 --- a/simple_syrup/runtime/settings_routes.py +++ b/simple_syrup/integration/settings_routes.py @@ -12,12 +12,12 @@ from typing import Any, Protocol, cast from aiohttp import web -from ..shared.logging import get_logger -from .settings import ( +from ..runtime.settings import ( SimpleSyrupSettings, SimpleSyrupSettingsError, ) -from .settings_repository import SimpleSyrupSettingsRepository +from ..runtime.settings_repository import SimpleSyrupSettingsRepository +from ..shared.logging import get_logger LOGGER = get_logger(__name__) SETTINGS_ROUTE = "/simple-syrup/settings" diff --git a/simple_syrup/masking/prompt_segs_with_sam_service.py b/simple_syrup/masking/prompt_segs_with_sam_service.py index c03a9a6..ae38b56 100644 --- a/simple_syrup/masking/prompt_segs_with_sam_service.py +++ b/simple_syrup/masking/prompt_segs_with_sam_service.py @@ -15,8 +15,7 @@ from ..domain.segs import ( NativeSegs, Segment, ) -from ..masking.mask_ops import MaskRefinementSettings, refine_prompt_mask -from ..masking.segs_mask_ops import ( +from ..domain.segs_mask_ops import ( crop_image, crop_mask, crop_region_for_bbox, @@ -24,6 +23,7 @@ from ..masking.segs_mask_ops import ( normalize_mask, validate_single_image, ) +from ..masking.mask_ops import MaskRefinementSettings, refine_prompt_mask from ..runtime.sam_segmenter import SAMBoxSegmenter, SAMModelSegmenter from ..runtime.text_box_detector import ( GroundingDINOTextBoxDetector, diff --git a/simple_syrup/masking/regional_detailing_masks.py b/simple_syrup/masking/regional_detailing_masks.py index e8e5071..2432fa8 100644 --- a/simple_syrup/masking/regional_detailing_masks.py +++ b/simple_syrup/masking/regional_detailing_masks.py @@ -21,8 +21,8 @@ from ..domain.regional_detailing import ( SegmentConditioningPair, ) from ..domain.segs import CropRegion +from ..domain.segs_mask_ops import feather_mask, resize_mask from .detailer_masks import gaussian_feather_mask -from .segs_mask_ops import feather_mask, resize_mask OPERATION = "Detail SEGS as Regions" diff --git a/simple_syrup/nodes/detailer_input_adapters.py b/simple_syrup/nodes/detailer_input_adapters.py index f348c44..d1ec835 100644 --- a/simple_syrup/nodes/detailer_input_adapters.py +++ b/simple_syrup/nodes/detailer_input_adapters.py @@ -12,7 +12,7 @@ import torch from ..domain.conditioning_batch import ConditioningBatch from ..domain.segs import NativeSegs -from ..masking.segs_mask_ops import iter_single_images, validate_image_batch +from ..domain.segs_mask_ops import iter_single_images, validate_image_batch def image_inputs(image: object, operation_name: str) -> tuple[torch.Tensor, ...]: diff --git a/simple_syrup/nodes/detect_segs_with_ultralytics.py b/simple_syrup/nodes/detect_segs_with_ultralytics.py index 188ce33..5466cec 100644 --- a/simple_syrup/nodes/detect_segs_with_ultralytics.py +++ b/simple_syrup/nodes/detect_segs_with_ultralytics.py @@ -16,8 +16,8 @@ from ..domain.segs import ( SORT_ORDER_OPTIONS, NativeSegs, ) -from ..masking.segs_mask_ops import iter_single_images, validate_image_batch -from ..runtime.ultralytics_loader import UltralyticsDetectorModel +from ..domain.segs_mask_ops import iter_single_images, validate_image_batch +from ..runtime.ultralytics_model_adapter import UltralyticsDetectorModel from ..services.segs_detection_service import ( SegsDetectionService, ) diff --git a/simple_syrup/nodes/load_ultralytics_model.py b/simple_syrup/nodes/load_ultralytics_model.py index 7810738..ad659f4 100644 --- a/simple_syrup/nodes/load_ultralytics_model.py +++ b/simple_syrup/nodes/load_ultralytics_model.py @@ -9,7 +9,7 @@ from __future__ import annotations from typing import Any, ClassVar from ..runtime.model_downloads import ComfyProgressReporter -from ..runtime.ultralytics_loader import UltralyticsLoaderService +from ..services.ultralytics_loader_service import UltralyticsLoaderService class LoadUltralyticsModel: diff --git a/simple_syrup/nodes/prompt_segs_with_sam.py b/simple_syrup/nodes/prompt_segs_with_sam.py index d009370..bcb8a83 100644 --- a/simple_syrup/nodes/prompt_segs_with_sam.py +++ b/simple_syrup/nodes/prompt_segs_with_sam.py @@ -16,9 +16,9 @@ from ..domain.segs import ( SORT_ORDER_OPTIONS, NativeSegs, ) +from ..domain.segs_mask_ops import iter_single_images, validate_image_batch from ..masking.mask_ops import DETAIL_METHODS from ..masking.prompt_segs_with_sam_service import PromptSEGSWithSAMService -from ..masking.segs_mask_ops import iter_single_images, validate_image_batch from ..services.segs_output_service import ( CombinedSegsResult, build_combined_segs_result, diff --git a/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py index 3166d73..9d79dc3 100644 --- a/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py +++ b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py @@ -8,7 +8,7 @@ from __future__ import annotations from typing import Any -from ..runtime.prompt_control_schedule_encode_graph import ( +from ..services.prompt_control_schedule_encode_graph import ( PromptControlScheduleEncodeGraphBuilder, ) diff --git a/simple_syrup/nodes/segs_from_sam_output.py b/simple_syrup/nodes/segs_from_sam_output.py index 2ec3e19..dc17684 100644 --- a/simple_syrup/nodes/segs_from_sam_output.py +++ b/simple_syrup/nodes/segs_from_sam_output.py @@ -11,7 +11,7 @@ from typing import Any, ClassVar import torch -from ..masking.segs_mask_ops import iter_single_images, validate_image_batch +from ..domain.segs_mask_ops import iter_single_images, validate_image_batch from ..runtime.progress import PhaseProgressReporter, create_comfy_phase_progress from ..services.segs_from_sam_output_service import SEGSFromSAMOutputService diff --git a/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py b/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py index bc64520..a00d055 100644 --- a/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py +++ b/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py @@ -10,7 +10,7 @@ from importlib import import_module from typing import TYPE_CHECKING, Any, ClassVar from ..domain.prompt_batch_parser import DEFAULT_PROMPT_BATCH_SEPARATOR -from ..runtime.prompt_control_batch_graph import PromptControlBatchGraphBuilder +from ..services.prompt_control_batch_graph import PromptControlBatchGraphBuilder if TYPE_CHECKING: diff --git a/simple_syrup/nodes_v3/mask_to_segs.py b/simple_syrup/nodes_v3/mask_to_segs.py index 45d80e6..fdbdea5 100644 --- a/simple_syrup/nodes_v3/mask_to_segs.py +++ b/simple_syrup/nodes_v3/mask_to_segs.py @@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, Any, ClassVar import torch from ..domain.segs import SORT_ORDER_OPTIONS, NativeSegs -from ..masking.segs_mask_ops import iter_single_images, validate_image_batch +from ..domain.segs_mask_ops import iter_single_images, validate_image_batch from ..services.mask_to_segs_service import MaskToSEGSService from ..services.segs_output_service import ( CombinedSegsResult, diff --git a/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py index 7ba6d70..3cef3f2 100644 --- a/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py +++ b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py @@ -9,7 +9,7 @@ from __future__ import annotations from importlib import import_module from typing import TYPE_CHECKING, Any, ClassVar -from ..runtime.prompt_control_schedule_encode_graph import ( +from ..services.prompt_control_schedule_encode_graph import ( PromptControlScheduleEncodeGraphBuilder, ) diff --git a/simple_syrup/runtime/comfy_conditioning_processing.py b/simple_syrup/runtime/comfy_conditioning_processing.py index fe4bfd7..da16895 100644 --- a/simple_syrup/runtime/comfy_conditioning_processing.py +++ b/simple_syrup/runtime/comfy_conditioning_processing.py @@ -13,6 +13,7 @@ from uuid import UUID import torch from comfy import sampler_helpers, samplers +from ..domain.attention_coupling_preparation import AttentionCouplingPreparation from ..domain.conditioning_schedule import ConditioningScheduleRange from ..domain.processed_regional_attention import ( ProcessedRegionalAttentionBranch, @@ -23,9 +24,6 @@ from ..domain.processed_regional_attention import ( from ..domain.raw_regional_attention import ( RawRegionalAttentionBranch, ) -from ..services.attention_coupling_preparation_service import ( - AttentionCouplingPreparation, -) from .attention_coupling.context_validation import RegionalContextValidator from .ppm_negpip_interop import PpmNegpipInterop diff --git a/simple_syrup/runtime/detail_preview_images.py b/simple_syrup/runtime/detail_preview_images.py new file mode 100644 index 0000000..05e6ff0 --- /dev/null +++ b/simple_syrup/runtime/detail_preview_images.py @@ -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 + +"""Convert and align image assets used by detail sampling previews.""" + +from __future__ import annotations + +import numpy as np +import torch +from PIL import Image + +from ..domain.segs import CropRegion + +CropBox = tuple[int, int, int, int] + + +def image_tensor_to_rgb_pil(image: torch.Tensor) -> Image.Image: + """Convert a single-image BHWC tensor to an RGB PIL image.""" + + if image.ndim != 4: + raise ValueError("detail preview image must be a BHWC tensor.") + if int(image.shape[0]) != 1: + raise ValueError("detail preview image must contain exactly one image.") + if int(image.shape[-1]) < 1: + raise ValueError("detail preview image must contain at least one channel.") + 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 normalize_mask_tensor(mask: torch.Tensor) -> torch.Tensor: + """Normalize an HW or single-item BHW mask tensor to HW float.""" + + working = mask.detach().float() + if working.ndim == 3 and int(working.shape[0]) == 1: + working = working[0] + if working.ndim != 2: + raise ValueError("detail preview work mask must be an HW tensor.") + return working + + +def detail_alpha_mask( + mask: torch.Tensor, + *, + source_size: tuple[int, int], + preview_size: tuple[int, int], + sampled_box: CropBox, + target_size: tuple[int, int], +) -> Image.Image: + """Return an alpha mask aligned to the sampled preview paste box.""" + + working = normalize_mask_tensor(mask).detach().cpu().clamp(0.0, 1.0) + mask_image = Image.fromarray((working.numpy() * 255.0).round().astype(np.uint8)) + if mask_image.size == source_size: + preview_mask = mask_image.resize(preview_size, Image.Resampling.BILINEAR) + return preview_mask.crop(sampled_box).resize( + target_size, + Image.Resampling.BILINEAR, + ) + return mask_image.resize(target_size, Image.Resampling.BILINEAR) + + +def validate_crop_region( + crop_region: CropRegion, + source_width: int, + source_height: int, +) -> None: + """Reject crop regions that cannot be mapped into the source image.""" + + if crop_region.left < 0 or crop_region.top < 0: + raise ValueError("crop_region left and top must be non-negative.") + if crop_region.right <= crop_region.left or crop_region.bottom <= crop_region.top: + raise ValueError("crop_region right/bottom must be greater than left/top.") + if crop_region.right > source_width or crop_region.bottom > source_height: + raise ValueError("crop_region must fit within the source image.") diff --git a/simple_syrup/runtime/detail_previews.py b/simple_syrup/runtime/detail_previews.py index c7890a4..ef2f0b5 100644 --- a/simple_syrup/runtime/detail_previews.py +++ b/simple_syrup/runtime/detail_previews.py @@ -15,6 +15,12 @@ import torch from PIL import Image, ImageDraw from ..domain.segs import CropRegion +from .detail_preview_images import ( + detail_alpha_mask, + image_tensor_to_rgb_pil, + normalize_mask_tensor, + validate_crop_region, +) DETAIL_PREVIEW_WASH_OPACITY = 0.55 DETAIL_PREVIEW_OUTLINE_RGB = (255, 0, 0) @@ -126,7 +132,7 @@ class DetailPreviewCompositor: ) -> DetailPreviewCompositor: """Create a compositor with detailer background work precomputed.""" - source_image = _image_tensor_to_rgb_pil(context.image) + source_image = image_tensor_to_rgb_pil(context.image) sampled_region = context.sampled_region or context.work_region geometry = build_detail_preview_geometry( source_width=source_image.width, @@ -147,7 +153,7 @@ class DetailPreviewCompositor: ) left, top, right, bottom = geometry.crop_box crop_size = (max(1, right - left), max(1, bottom - top)) - detail_alpha_mask = _detail_alpha_mask( + detail_alpha = detail_alpha_mask( context.work_mask, source_size=(source_image.width, source_image.height), preview_size=geometry.preview_size, @@ -157,7 +163,7 @@ class DetailPreviewCompositor: return cls( geometry=geometry, washed_background=washed_background, - detail_alpha_mask=detail_alpha_mask, + detail_alpha_mask=detail_alpha, ) def compose(self, crop_preview: Image.Image) -> Image.Image: @@ -211,9 +217,9 @@ def build_detail_preview_geometry( source_height, max_preview_resolution, ) - _validate_crop_region(crop_region, source_width, source_height) + validate_crop_region(crop_region, source_width, source_height) resolved_outline_region = outline_region or crop_region - _validate_crop_region(resolved_outline_region, source_width, source_height) + validate_crop_region(resolved_outline_region, source_width, source_height) preview_width, preview_height = preview_size scale_x = float(preview_width) / float(source_width) @@ -238,7 +244,7 @@ def build_detail_preview_geometry( def work_region_from_mask(mask: torch.Tensor) -> CropRegion: """Return the tight work region around a non-empty HW or single-item BHW mask.""" - working = _normalize_mask_tensor(mask) + working = normalize_mask_tensor(mask) coordinates = torch.nonzero(working > 0, as_tuple=False) if coordinates.numel() == 0: raise ValueError("detail preview work mask must contain at least one pixel.") @@ -367,83 +373,6 @@ def _source_outline_box( ) -def _image_tensor_to_rgb_pil(image: torch.Tensor) -> Image.Image: - """Convert a single-image BHWC tensor to an RGB PIL image.""" - - import numpy as np - - if image.ndim != 4: - raise ValueError("detail preview image must be a BHWC tensor.") - if int(image.shape[0]) != 1: - raise ValueError("detail preview image must contain exactly one image.") - if int(image.shape[-1]) < 1: - raise ValueError("detail preview image must contain at least one channel.") - - 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 _mask_tensor_to_l_pil(mask: torch.Tensor) -> Image.Image: - """Convert an HW or single-item BHW mask tensor to a grayscale alpha image.""" - - import numpy as np - - working = _normalize_mask_tensor(mask).detach().cpu().clamp(0.0, 1.0) - return Image.fromarray((working.numpy() * 255.0).round().astype(np.uint8)) - - -def _normalize_mask_tensor(mask: torch.Tensor) -> torch.Tensor: - """Normalize an HW or single-item BHW mask tensor to HW float.""" - - working = mask.detach().float() - if working.ndim == 3 and int(working.shape[0]) == 1: - working = working[0] - if working.ndim != 2: - raise ValueError("detail preview work mask must be an HW tensor.") - return working - - -def _detail_alpha_mask( - mask: torch.Tensor, - *, - source_size: tuple[int, int], - preview_size: tuple[int, int], - sampled_box: CropBox, - target_size: tuple[int, int], -) -> Image.Image: - """Return an alpha mask aligned to the sampled preview paste box.""" - - mask_image = _mask_tensor_to_l_pil(mask) - if mask_image.size == source_size: - preview_mask = mask_image.resize(preview_size, Image.Resampling.BILINEAR) - return preview_mask.crop(sampled_box).resize( - target_size, - Image.Resampling.BILINEAR, - ) - return mask_image.resize(target_size, Image.Resampling.BILINEAR) - - -def _validate_crop_region( - crop_region: CropRegion, - source_width: int, - source_height: int, -) -> None: - """Reject crop regions that cannot be mapped into the source image.""" - - if crop_region.left < 0 or crop_region.top < 0: - raise ValueError("crop_region left and top must be non-negative.") - if crop_region.right <= crop_region.left or crop_region.bottom <= crop_region.top: - raise ValueError("crop_region right/bottom must be greater than left/top.") - if crop_region.right > source_width or crop_region.bottom > source_height: - raise ValueError("crop_region must fit within the source image.") - - def _validate_positive_int(name: str, value: int) -> None: """Reject non-positive integer values.""" diff --git a/simple_syrup/runtime/external_llm_images.py b/simple_syrup/runtime/external_llm_images.py index 797aae9..2b8ee26 100644 --- a/simple_syrup/runtime/external_llm_images.py +++ b/simple_syrup/runtime/external_llm_images.py @@ -13,7 +13,7 @@ import torch from PIL import Image from ..domain.segs import Segment -from ..masking.segs_mask_ops import crop_image, crop_mask, resize_mask +from ..domain.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") diff --git a/simple_syrup/runtime/regional_detail_sampler.py b/simple_syrup/runtime/regional_detail_sampler.py new file mode 100644 index 0000000..b8113ac --- /dev/null +++ b/simple_syrup/runtime/regional_detail_sampler.py @@ -0,0 +1,72 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Adapt shared detail sampling to regional MultiDiffusion runtime calls.""" + +from __future__ import annotations + +from typing import Any + +import torch + +from ..domain.regional_detailing import LatentRegion +from . import regional_multidiffusion_sampling +from .detail_previews import DetailPreviewContext +from .detail_sampling import DetailSampler, Latent + + +class RegionalDetailSampler: + """Adapt shared detail sampling helpers to regional MultiDiffusion.""" + + def __init__(self, detail_sampler: DetailSampler | None = None) -> None: + """Create the runtime adapter with injectable encode/decode behavior.""" + + self._detail_sampler = detail_sampler or DetailSampler() + + def encode(self, vae: Any, pixels: torch.Tensor, tiled: bool) -> Latent: + """Encode pixels into a latent dictionary.""" + + return self._detail_sampler.encode(vae, pixels, tiled) + + def decode(self, vae: Any, latent: Latent, tiled: bool) -> torch.Tensor: + """Decode latent samples into pixels.""" + + return self._detail_sampler.decode(vae, latent, tiled) + + def sample_regions( + self, + *, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + regions: tuple[LatentRegion, ...], + denoise: float, + global_prompt_weight: float, + preview_context: DetailPreviewContext | None = None, + differential_diffusion: bool = False, + ) -> Latent: + """Sample one full latent with regional MultiDiffusion.""" + + return regional_multidiffusion_sampling.sample_regional_multidiffusion( + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, + regions=regions, + denoise=denoise, + global_prompt_weight=global_prompt_weight, + preview_context=preview_context, + differential_diffusion=differential_diffusion, + ) diff --git a/simple_syrup/runtime/regional_multidiffusion_prediction.py b/simple_syrup/runtime/regional_multidiffusion_prediction.py new file mode 100644 index 0000000..bd2f821 --- /dev/null +++ b/simple_syrup/runtime/regional_multidiffusion_prediction.py @@ -0,0 +1,293 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Blend regional predictions inside ComfyUI calc-cond-batch execution.""" + +from __future__ import annotations + +from collections.abc import Callable +from importlib import import_module +from types import ModuleType +from typing import Any, TypeAlias, cast + +import torch + +from ..domain.regional_detailing import LatentRegion +from .tiled_sampling_validation import validate_tensor_shape + +SAMPLER_LABEL = "Regional MultiDiffusion" +CalcCondBatchFunction: TypeAlias = Callable[[dict[str, Any]], list[torch.Tensor]] + + +class RegionalMultiDiffusionCalcCondBatch: + """Blend regional condition predictions before CFG is applied.""" + + def __init__( + self, + *, + latent_width: int, + latent_height: int, + regions: tuple[LatentRegion, ...], + existing_calc_cond_batch: CalcCondBatchFunction | None, + global_prompt_weight: float, + ) -> None: + """Create a calc-cond-batch wrapper for one latent sampling shape.""" + + self._latent_width = latent_width + self._latent_height = latent_height + self._regions = regions + self._existing_calc_cond_batch = existing_calc_cond_batch + self._global_prompt_weight = global_prompt_weight + + def __call__(self, args: dict[str, Any]) -> list[torch.Tensor]: + """Return fallback predictions blended with regional predictions.""" + + x = args["input"] + if not isinstance(x, torch.Tensor): + raise ValueError("Regional MultiDiffusion model input must be a tensor.") + validate_tensor_shape(x, sampler_label=SAMPLER_LABEL) + if x.shape[-2:] != (self._latent_height, self._latent_width): + return self._call_original(args) + if not self._regions: + return self._call_original(args) + + timestep = args["sigma"] + if not isinstance(timestep, torch.Tensor): + raise ValueError("Regional MultiDiffusion sigma must be a tensor.") + conds = args["conds"] + if not isinstance(conds, list) or not conds: + raise ValueError("Regional MultiDiffusion conds must be a non-empty list.") + + fallback = self._call_original(args) + regional_buffers = [torch.zeros_like(output) for output in fallback] + regional_weights = [_new_spatial_weight(x) for _output in fallback] + input_batch_size = int(x.shape[0]) + + for region in self._regions: + region_slice = _region_slicer(region, x.ndim) + region_x = x[region_slice] + region_conds = [ + _prepare_region_conditioning(region.positive, args=args, x=region_x), + *conds[1:], + ] + region_args = args.copy() + region_args["conds"] = region_conds + region_args["input"] = region_x + region_args["sigma"] = timestep + region_outputs = self._call_original(region_args) + self._accumulate_region_outputs( + outputs=region_outputs, + buffers=regional_buffers, + weights=regional_weights, + region=region, + input_batch_size=input_batch_size, + ) + + return [ + _blend_prediction( + fallback_output, + region_output, + region_weight, + global_prompt_weight=self._global_prompt_weight, + ) + for fallback_output, region_output, region_weight in zip( + fallback, + regional_buffers, + regional_weights, + strict=True, + ) + ] + + def _call_original(self, args: dict[str, Any]) -> list[torch.Tensor]: + """Call the previous calc-cond-batch hook or ComfyUI default.""" + + clean_args = args.copy() + clean_options = _clean_model_options( + cast(dict[str, Any], clean_args["model_options"]), + self._existing_calc_cond_batch, + ) + clean_args["model_options"] = clean_options + if self._existing_calc_cond_batch is not None: + return self._existing_calc_cond_batch(clean_args) + comfy_samplers = _comfy_samplers() + return cast( + list[torch.Tensor], + comfy_samplers.calc_cond_batch( + clean_args["model"], + clean_args["conds"], + clean_args["input"], + clean_args["sigma"], + clean_options, + ), + ) + + def _accumulate_region_outputs( + self, + *, + outputs: list[torch.Tensor], + buffers: list[torch.Tensor], + weights: list[torch.Tensor], + region: LatentRegion, + input_batch_size: int, + ) -> None: + """Accumulate one region prediction into full-latent buffers.""" + + for output_index, output in enumerate(outputs[: len(buffers)]): + region_slice = _region_slicer(region, output.ndim) + box = region.latent_box + mask_slice = ( + region.latent_mask[ + box.y : box.y + box.height, + box.x : box.x + box.width, + ] + .reshape((1,) * (output.ndim - 2) + (box.height, box.width)) + .to(device=output.device, dtype=torch.float32) + ) + buffers[output_index][region_slice] += output[ + :input_batch_size + ] * mask_slice.to(dtype=output.dtype) + weights[output_index][region_slice] += mask_slice + + +def validate_regions( + *, + latent_width: int, + latent_height: int, + regions: tuple[LatentRegion, ...], +) -> None: + """Reject regions incompatible with the current latent shape.""" + + for region in regions: + box = region.latent_box + if region.latent_mask.shape != (latent_height, latent_width): + raise ValueError( + f"Regional MultiDiffusion region {region.index} ('{region.label}') " + "latent_mask must match the full latent height and width." + ) + if box.x < 0 or box.y < 0 or box.width < 1 or box.height < 1: + raise ValueError( + f"Regional MultiDiffusion region {region.index} ('{region.label}') " + "has an invalid latent box." + ) + if box.x + box.width > latent_width or box.y + box.height > latent_height: + raise ValueError( + f"Regional MultiDiffusion region {region.index} ('{region.label}') " + "latent box must fit inside the latent." + ) + + +def _region_slicer(region: LatentRegion, tensor_ndim: int) -> tuple[slice, ...]: + """Return a slicer that crops a tensor to a latent region box.""" + + box = region.latent_box + return ( + (slice(None),) * (tensor_ndim - 2) + + (slice(box.y, box.y + box.height),) + + (slice(box.x, box.x + box.width),) + ) + + +def _new_spatial_weight(x: torch.Tensor) -> torch.Tensor: + """Create a full-latent spatial weight buffer.""" + + return torch.zeros( + (1,) * (x.ndim - 2) + (int(x.shape[-2]), int(x.shape[-1])), + device=x.device, + dtype=torch.float32, + ) + + +def _blend_prediction( + fallback: torch.Tensor, + regional: torch.Tensor, + weight: torch.Tensor, + *, + global_prompt_weight: float, +) -> torch.Tensor: + """Blend normalized regional predictions with global fallback predictions.""" + + has_region = weight > 0 + normalized = torch.where( + has_region, + regional / torch.clamp(weight, min=1.0e-37).to(dtype=regional.dtype), + regional, + ) + coverage = torch.clamp(weight, 0.0, 1.0).to(dtype=fallback.dtype) + regional_alpha = coverage * (1.0 - global_prompt_weight) + blended = fallback * (1.0 - regional_alpha) + normalized * regional_alpha + return torch.where(has_region, blended, fallback) + + +def _prepare_region_conditioning( + conditioning: object, + *, + args: dict[str, Any], + x: torch.Tensor, +) -> list[dict[str, Any]]: + """Convert raw Comfy CONDITIONING into sampler-ready dictionaries.""" + + if _is_processed_conditioning(conditioning): + return cast(list[dict[str, Any]], conditioning) + if not isinstance(conditioning, list): + raise TypeError("Regional MultiDiffusion region_positive must be CONDITIONING.") + sampler_helpers = _comfy_sampler_helpers() + comfy_samplers = _comfy_samplers() + model = args["model"] + converted = cast(list[dict[str, Any]], sampler_helpers.convert_cond(conditioning)) + comfy_samplers.resolve_areas_and_cond_masks_multidim( + converted, + tuple(int(dim) for dim in x.shape[2:]), + x.device, + ) + comfy_samplers.calculate_start_end_timesteps(model, converted) + if hasattr(model, "extra_conds"): + converted = cast( + list[dict[str, Any]], + comfy_samplers.encode_model_conds( + model.extra_conds, + converted, + x, + x.device, + "positive", + ), + ) + return converted + + +def _is_processed_conditioning(conditioning: object) -> bool: + """Return whether a conditioning value is already sampler-ready.""" + + if not isinstance(conditioning, list): + return False + if not conditioning: + return True + return all( + isinstance(item, dict) and "model_conds" in item for item in conditioning + ) + + +def _clean_model_options( + model_options: dict[str, Any], + existing_calc_cond_batch: CalcCondBatchFunction | None, +) -> dict[str, Any]: + """Return model options that cannot recurse into this wrapper.""" + + clean_options = model_options.copy() + if existing_calc_cond_batch is None: + clean_options.pop("sampler_calc_cond_batch_function", None) + else: + clean_options["sampler_calc_cond_batch_function"] = existing_calc_cond_batch + return clean_options + + +def _comfy_samplers() -> ModuleType: + """Import Comfy sampler helpers lazily.""" + + return import_module("comfy.samplers") + + +def _comfy_sampler_helpers() -> ModuleType: + """Import Comfy sampler conditioning helpers lazily.""" + + return import_module("comfy.sampler_helpers") diff --git a/simple_syrup/runtime/regional_multidiffusion_sampling.py b/simple_syrup/runtime/regional_multidiffusion_sampling.py index cacf701..bdb2328 100644 --- a/simple_syrup/runtime/regional_multidiffusion_sampling.py +++ b/simple_syrup/runtime/regional_multidiffusion_sampling.py @@ -10,13 +10,10 @@ from __future__ import annotations -from collections.abc import Callable from dataclasses import dataclass from importlib import import_module from types import ModuleType -from typing import Any, TypeAlias, cast - -import torch +from typing import Any, cast from ..domain.regional_detailing import LatentRegion from ..domain.regional_features import EMPTY_REGIONAL_CAPABILITY_ADMISSION @@ -29,6 +26,11 @@ from .differential_diffusion import ( ) from .model_patcher_mutations import ModelCalcCondBatchMutation from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelMutation +from .regional_multidiffusion_prediction import ( + CalcCondBatchFunction, + RegionalMultiDiffusionCalcCondBatch, + validate_regions, +) from .tiled_sampling_validation import ( Latent, reject_unsupported_conditioning, @@ -39,7 +41,6 @@ from .tiled_sampling_validation import ( LOGGER = get_logger(__name__) SAMPLER_LABEL = "Regional MultiDiffusion" UNIPC_SAMPLERS = frozenset({"uni_pc", "uni_pc_bh2"}) -CalcCondBatchFunction: TypeAlias = Callable[[dict[str, Any]], list[torch.Tensor]] @dataclass(frozen=True) @@ -193,7 +194,7 @@ def clone_model_with_regional_multidiffusion( denoise=1.0, global_prompt_weight=global_prompt_weight, ) - _validate_regions( + validate_regions( latent_width=latent_width, latent_height=latent_height, regions=regions, @@ -234,220 +235,6 @@ def clone_model_with_regional_multidiffusion( return derived_model, summary -class RegionalMultiDiffusionCalcCondBatch: - """Blend regional condition predictions before CFG is applied.""" - - def __init__( - self, - *, - latent_width: int, - latent_height: int, - regions: tuple[LatentRegion, ...], - existing_calc_cond_batch: CalcCondBatchFunction | None, - global_prompt_weight: float, - ) -> None: - """Create a calc-cond-batch wrapper for one latent sampling shape.""" - - self._latent_width = latent_width - self._latent_height = latent_height - self._regions = regions - self._existing_calc_cond_batch = existing_calc_cond_batch - self._global_prompt_weight = global_prompt_weight - - def __call__(self, args: dict[str, Any]) -> list[torch.Tensor]: - """Return fallback predictions blended with regional predictions.""" - - x = args["input"] - if not isinstance(x, torch.Tensor): - raise ValueError("Regional MultiDiffusion model input must be a tensor.") - validate_tensor_shape(x, sampler_label=SAMPLER_LABEL) - if x.shape[-2:] != (self._latent_height, self._latent_width): - return self._call_original(args) - if not self._regions: - return self._call_original(args) - - timestep = args["sigma"] - if not isinstance(timestep, torch.Tensor): - raise ValueError("Regional MultiDiffusion sigma must be a tensor.") - conds = args["conds"] - if not isinstance(conds, list) or not conds: - raise ValueError("Regional MultiDiffusion conds must be a non-empty list.") - - fallback = self._call_original(args) - regional_buffers = [torch.zeros_like(output) for output in fallback] - regional_weights = [_new_spatial_weight(x) for _output in fallback] - input_batch_size = int(x.shape[0]) - - for region in self._regions: - region_slice = _region_slicer(region, x.ndim) - region_x = x[region_slice] - region_conds = [ - _prepare_region_conditioning( - region.positive, - args=args, - x=region_x, - ), - *conds[1:], - ] - region_args = args.copy() - region_args["conds"] = region_conds - region_args["input"] = region_x - region_args["sigma"] = timestep - region_outputs = self._call_original(region_args) - self._accumulate_region_outputs( - outputs=region_outputs, - buffers=regional_buffers, - weights=regional_weights, - region=region, - input_batch_size=input_batch_size, - ) - - return [ - _blend_prediction( - fallback_output, - region_output, - region_weight, - global_prompt_weight=self._global_prompt_weight, - ) - for fallback_output, region_output, region_weight in zip( - fallback, - regional_buffers, - regional_weights, - strict=True, - ) - ] - - def _call_original(self, args: dict[str, Any]) -> list[torch.Tensor]: - """Call the previous calc-cond-batch hook or ComfyUI default.""" - - clean_args = args.copy() - clean_options = _clean_model_options( - cast(dict[str, Any], clean_args["model_options"]), - self._existing_calc_cond_batch, - ) - clean_args["model_options"] = clean_options - if self._existing_calc_cond_batch is not None: - return self._existing_calc_cond_batch(clean_args) - comfy_samplers = _comfy_samplers() - return cast( - list[torch.Tensor], - comfy_samplers.calc_cond_batch( - clean_args["model"], - clean_args["conds"], - clean_args["input"], - clean_args["sigma"], - clean_options, - ), - ) - - def _accumulate_region_outputs( - self, - *, - outputs: list[torch.Tensor], - buffers: list[torch.Tensor], - weights: list[torch.Tensor], - region: LatentRegion, - input_batch_size: int, - ) -> None: - """Accumulate one region prediction into full-latent buffers.""" - - for output_index, output in enumerate(outputs[: len(buffers)]): - region_slice = _region_slicer(region, output.ndim) - box = region.latent_box - mask_slice = ( - region.latent_mask[ - box.y : box.y + box.height, - box.x : box.x + box.width, - ] - .reshape((1,) * (output.ndim - 2) + (box.height, box.width)) - .to(device=output.device, dtype=torch.float32) - ) - buffers[output_index][region_slice] += output[ - :input_batch_size - ] * mask_slice.to(dtype=output.dtype) - weights[output_index][region_slice] += mask_slice - - -def _validate_regions( - *, - latent_width: int, - latent_height: int, - regions: tuple[LatentRegion, ...], -) -> None: - """Reject regions incompatible with the current latent shape.""" - - for region in regions: - _validate_region(region, latent_width=latent_width, latent_height=latent_height) - - -def _validate_region( - region: LatentRegion, - *, - latent_width: int, - latent_height: int, -) -> None: - """Reject regions incompatible with the current latent shape.""" - - box = region.latent_box - if region.latent_mask.shape != (latent_height, latent_width): - raise ValueError( - f"Regional MultiDiffusion region {region.index} ('{region.label}') " - "latent_mask must match the full latent height and width." - ) - if box.x < 0 or box.y < 0 or box.width < 1 or box.height < 1: - raise ValueError( - f"Regional MultiDiffusion region {region.index} ('{region.label}') " - "has an invalid latent box." - ) - if box.x + box.width > latent_width or box.y + box.height > latent_height: - raise ValueError( - f"Regional MultiDiffusion region {region.index} ('{region.label}') " - "latent box must fit inside the latent." - ) - - -def _region_slicer(region: LatentRegion, tensor_ndim: int) -> tuple[slice, ...]: - """Return a slicer that crops a tensor to a latent region box.""" - - box = region.latent_box - return ( - (slice(None),) * (tensor_ndim - 2) - + (slice(box.y, box.y + box.height),) - + (slice(box.x, box.x + box.width),) - ) - - -def _new_spatial_weight(x: torch.Tensor) -> torch.Tensor: - """Create a full-latent spatial weight buffer.""" - - return torch.zeros( - (1,) * (x.ndim - 2) + (int(x.shape[-2]), int(x.shape[-1])), - device=x.device, - dtype=torch.float32, - ) - - -def _blend_prediction( - fallback: torch.Tensor, - regional: torch.Tensor, - weight: torch.Tensor, - *, - global_prompt_weight: float, -) -> torch.Tensor: - """Blend normalized regional predictions with global fallback predictions.""" - - has_region = weight > 0 - normalized = torch.where( - has_region, - regional / torch.clamp(weight, min=1.0e-37).to(dtype=regional.dtype), - regional, - ) - coverage = torch.clamp(weight, 0.0, 1.0).to(dtype=fallback.dtype) - regional_alpha = coverage * (1.0 - global_prompt_weight) - blended = fallback * (1.0 - regional_alpha) + normalized * regional_alpha - return torch.where(has_region, blended, fallback) - - def _validate_sampling_controls( *, steps: int, @@ -470,69 +257,6 @@ def _validate_global_prompt_weight(global_prompt_weight: float) -> None: raise ValueError("global_prompt_weight must be between 0.0 and 1.0.") -def _prepare_region_conditioning( - conditioning: object, - *, - args: dict[str, Any], - x: torch.Tensor, -) -> list[dict[str, Any]]: - """Convert raw Comfy CONDITIONING into sampler-ready condition dictionaries.""" - - if _is_processed_conditioning(conditioning): - return cast(list[dict[str, Any]], conditioning) - if not isinstance(conditioning, list): - raise TypeError("Regional MultiDiffusion region_positive must be CONDITIONING.") - - sampler_helpers = _comfy_sampler_helpers() - comfy_samplers = _comfy_samplers() - model = args["model"] - converted = cast(list[dict[str, Any]], sampler_helpers.convert_cond(conditioning)) - comfy_samplers.resolve_areas_and_cond_masks_multidim( - converted, - tuple(int(dim) for dim in x.shape[2:]), - x.device, - ) - comfy_samplers.calculate_start_end_timesteps(model, converted) - if hasattr(model, "extra_conds"): - converted = cast( - list[dict[str, Any]], - comfy_samplers.encode_model_conds( - model.extra_conds, - converted, - x, - x.device, - "positive", - ), - ) - return converted - - -def _is_processed_conditioning(conditioning: object) -> bool: - """Return whether a conditioning value is already sampler-ready.""" - - if not isinstance(conditioning, list): - return False - if not conditioning: - return True - return all( - isinstance(item, dict) and "model_conds" in item for item in conditioning - ) - - -def _clean_model_options( - model_options: dict[str, Any], - existing_calc_cond_batch: CalcCondBatchFunction | None, -) -> dict[str, Any]: - """Return model options that cannot recurse into this wrapper.""" - - clean_options = model_options.copy() - if existing_calc_cond_batch is None: - clean_options.pop("sampler_calc_cond_batch_function", None) - else: - clean_options["sampler_calc_cond_batch_function"] = existing_calc_cond_batch - return clean_options - - def _reject_unipc_sampler(sampler_name: str) -> None: """Reject UniPC samplers because MultiDiffusion is incompatible with them.""" @@ -566,18 +290,6 @@ def _comfy_utils() -> ModuleType: return import_module("comfy.utils") -def _comfy_samplers() -> ModuleType: - """Import ComfyUI sampler helpers lazily.""" - - return import_module("comfy.samplers") - - -def _comfy_sampler_helpers() -> ModuleType: - """Import ComfyUI conditioning conversion helpers lazily.""" - - return import_module("comfy.sampler_helpers") - - def _latent_preview() -> ModuleType: """Import ComfyUI preview helpers lazily.""" diff --git a/simple_syrup/runtime/seg_preview_assets.py b/simple_syrup/runtime/seg_preview_assets.py index 9785889..4ad13eb 100644 --- a/simple_syrup/runtime/seg_preview_assets.py +++ b/simple_syrup/runtime/seg_preview_assets.py @@ -12,7 +12,7 @@ from typing import Any, Protocol import torch -from ..services.simple_preview_segs_service import SegPreviewDocument +from ..domain.seg_preview import SegPreviewDocument SEG_PREVIEW_UI_KEY = "simple_syrup_segs_preview" diff --git a/simple_syrup/runtime/ultralytics_detection.py b/simple_syrup/runtime/ultralytics_detection.py index 3e1e811..a8ccd6b 100644 --- a/simple_syrup/runtime/ultralytics_detection.py +++ b/simple_syrup/runtime/ultralytics_detection.py @@ -12,8 +12,8 @@ from dataclasses import dataclass import torch from ..domain.segs import BoundingBox -from ..masking.segs_mask_ops import normalize_mask -from .ultralytics_loader import UltralyticsDetectorModel +from ..domain.segs_mask_ops import normalize_mask +from .ultralytics_model_adapter import UltralyticsDetectorModel @dataclass(frozen=True) diff --git a/simple_syrup/runtime/ultralytics_loader.py b/simple_syrup/runtime/ultralytics_loader.py deleted file mode 100644 index 183ea3a..0000000 --- a/simple_syrup/runtime/ultralytics_loader.py +++ /dev/null @@ -1,544 +0,0 @@ -# SimpleSyrup - workflow-focused ComfyUI extensions for image generation -# Copyright (C) 2026 Artificial Sweetener and contributors -# SPDX-License-Identifier: AGPL-3.0-or-later - -"""Ultralytics detector model discovery and lazy loading.""" - -from __future__ import annotations - -import importlib -from collections.abc import MutableMapping -from dataclasses import dataclass -from pathlib import Path -from types import ModuleType -from typing import Any, TypeAlias, cast - -from ..shared.logging import get_logger -from .model_catalog import ULTRALYTICS_ENTRIES, ModelEntry -from .model_choices import ModelChoiceService -from .model_downloads import DownloadRequest, ModelDownloader, ProgressReporter -from .model_folders import ( - SUPPORTED_MODEL_EXTENSIONS, - expected_model_file, - resolve_model_file, -) -from .model_instance_cache import ModelInstanceCache - -LOGGER = get_logger(__name__) - -NO_LOCAL_ULTRALYTICS_MODELS = "No local Ultralytics models found" -ULTRALYTICS_FOLDER = "ultralytics" -ULTRALYTICS_BBOX_FOLDER = "ultralytics_bbox" -ULTRALYTICS_SEGM_FOLDER = "ultralytics_segm" - -ModelFolderRegistry: TypeAlias = dict[str, tuple[list[str], set[str]]] - - -@dataclass(frozen=True) -class UltralyticsDetectorModel: - """Store a loaded Ultralytics detector with SimpleSyrup metadata.""" - - model_name: str - model_path: Path - model: Any - task: str - names: dict[int, str] - supports_segmentation: bool - - -@dataclass(frozen=True) -class LoadedUltralyticsDetector: - """Bundle native and compatibility detector outputs from the loader.""" - - detector_model: UltralyticsDetectorModel - bbox_detector: object - segm_detector: object - - -@dataclass(frozen=True) -class UltralyticsModelCacheKey: - """Identify a loaded Ultralytics detector for process-level reuse.""" - - model_path: Path - - -_LOADED_ULTRALYTICS_MODELS: dict[ - UltralyticsModelCacheKey, LoadedUltralyticsDetector -] = {} - - -class UltralyticsLoaderService: - """Discover and load Ultralytics detector models from ComfyUI folders.""" - - def __init__( - self, - folder_paths_module: ModuleType | None = None, - ultralytics_module: ModuleType | None = None, - downloader: ModelDownloader | None = None, - choice_service: ModelChoiceService | None = None, - cache: ( - MutableMapping[UltralyticsModelCacheKey, LoadedUltralyticsDetector] | None - ) = None, - ) -> None: - """Create the loader with injectable runtime modules for tests.""" - - self._folder_paths_module = folder_paths_module - self._ultralytics_module = ultralytics_module - self._downloader = downloader or ModelDownloader() - self._choice_service = choice_service or ModelChoiceService() - self._cache: ModelInstanceCache[ - UltralyticsModelCacheKey, LoadedUltralyticsDetector - ] = ModelInstanceCache( - cache if cache is not None else _LOADED_ULTRALYTICS_MODELS - ) - - def model_choices(self) -> list[str]: - """Return installed choices first, followed by downloadable catalog choices.""" - - self._register_model_folders() - curated_choices = self._choice_service.ultralytics_choices() - catalog_choice_labels = { - _catalog_selection(entry): entry.display_name - for entry in ULTRALYTICS_ENTRIES - } - available_choices = self.available_models() - visible_catalog_choices = set(curated_choices) - installed_catalog_choices = [ - entry.display_name - for entry in ULTRALYTICS_ENTRIES - if ( - entry.display_name in visible_catalog_choices - and _catalog_selection(entry) in available_choices - ) - ] - installed_non_catalog_choices = [ - choice - for choice in available_choices - if choice not in catalog_choice_labels - ] - downloadable_choices = [ - choice - for choice in curated_choices - if choice not in installed_catalog_choices - ] - choices = ( - installed_non_catalog_choices - + installed_catalog_choices - + downloadable_choices - ) - return choices or [NO_LOCAL_ULTRALYTICS_MODELS] - - def available_models(self) -> list[str]: - """Return supported model files in registered Ultralytics folders.""" - - self._register_model_folders() - folder_paths = self._folder_paths() - choices: set[str] = set() - - for folder in self._folder_paths_for(ULTRALYTICS_FOLDER): - if not folder.is_dir(): - continue - choices.update(path.name for path in _supported_files(folder)) - bbox_dir = folder / "bbox" - segm_dir = folder / "segm" - choices.update(f"bbox/{path.name}" for path in _supported_files(bbox_dir)) - choices.update(f"segm/{path.name}" for path in _supported_files(segm_dir)) - - for path in self._folder_paths_for(ULTRALYTICS_BBOX_FOLDER): - choices.update(f"bbox/{file.name}" for file in _supported_files(path)) - for path in self._folder_paths_for(ULTRALYTICS_SEGM_FOLDER): - choices.update(f"segm/{file.name}" for file in _supported_files(path)) - - registry = cast( - ModelFolderRegistry, - getattr(folder_paths, "folder_names_and_paths", {}), - ) - for folder_name in ( - ULTRALYTICS_FOLDER, - ULTRALYTICS_BBOX_FOLDER, - ULTRALYTICS_SEGM_FOLDER, - ): - if folder_name not in registry: - continue - for filename in folder_paths.get_filename_list(folder_name): - path = Path(str(filename)) - if path.suffix.lower() not in SUPPORTED_MODEL_EXTENSIONS: - continue - if folder_name == ULTRALYTICS_BBOX_FOLDER: - choices.add(f"bbox/{path.name}") - elif folder_name == ULTRALYTICS_SEGM_FOLDER: - choices.add(f"segm/{path.name}") - else: - choices.add(path.as_posix()) - - return sorted(choices) - - def load( - self, - model_name: str, - progress: ProgressReporter | None = None, - ) -> LoadedUltralyticsDetector: - """Load one Ultralytics model and create compatibility facades.""" - - self.reject_sentinel(model_name) - entry = _catalog_entry_or_none(model_name) - if entry is None: - model_path = self.resolve_model_path(model_name) - normalized_name = _normalized_model_name(model_name) - else: - model_path = self._resolve_catalog_entry(entry, progress) - normalized_name = _catalog_selection(entry) - - key = UltralyticsModelCacheKey(model_path=model_path.resolve()) - already_loaded = key in self._cache.entries - loaded = self._cache.get_or_load( - key, - lambda: self._load_uncached_detector(normalized_name, model_path), - ) - if already_loaded: - LOGGER.info( - "Ultralytics model loaded from process cache", - extra={ - "operation": "load_ultralytics_model", - "model_name": normalized_name, - "model_path": str(model_path), - "task": loaded.detector_model.task, - }, - ) - return loaded - - def _resolve_catalog_entry( - self, - entry: ModelEntry, - progress: ProgressReporter | None, - ) -> Path: - """Resolve or securely download one curated Ultralytics checkpoint.""" - - if len(entry.artifacts) != 1: - raise RuntimeError( - f"Ultralytics catalog entry '{entry.entry_id}' must have one artifact." - ) - - self._register_model_folders() - artifact = entry.artifacts[0] - existing = resolve_model_file( - artifact.folder_name, - artifact.filename, - self._folder_paths_module, - ) - destination = existing or expected_model_file( - artifact.folder_name, artifact.filename, self._folder_paths_module - ) - result = self._downloader.download( - DownloadRequest( - source_url=artifact.source_url, - destination_path=destination, - expected_folder=destination.parent, - description=artifact.description, - expected_sha256=artifact.sha256, - ), - progress, - ) - return result.path - - def _load_uncached_detector( - self, - model_name: str, - model_path: Path, - ) -> LoadedUltralyticsDetector: - """Load an Ultralytics detector after path resolution and cache lookup.""" - - ultralytics_module = self._ultralytics() - model_class = getattr(ultralytics_module, "YOLO", None) - if model_class is None: - raise RuntimeError( - "Ultralytics support requires a module exposing the YOLO class." - ) - - try: - raw_model = model_class(str(model_path)) - except Exception as exc: - LOGGER.error( - "Failed to load Ultralytics model", - extra={ - "operation": "load_ultralytics_model", - "model_name": model_name, - "model_path": str(model_path), - }, - exc_info=True, - ) - raise RuntimeError( - f"Ultralytics model '{model_name}' could not be loaded from " - f"'{model_path}'." - ) from exc - - task = _model_task(model_name, raw_model) - device_hint = _model_device_hint(raw_model) - detector_model = UltralyticsDetectorModel( - model_name=model_name, - model_path=model_path, - model=raw_model, - task=task, - names=_model_names(raw_model), - supports_segmentation=task in {"segment", "segm"}, - ) - - from .detector_compat import BBoxDetectorFacade, SegmDetectorFacade - - bbox_detector = BBoxDetectorFacade(detector_model) - segm_detector = SegmDetectorFacade(detector_model, bbox_detector) - LOGGER.info( - "Ultralytics model loaded", - extra={ - "operation": "load_ultralytics_model", - "model_name": model_name, - "model_path": str(model_path), - "task": task, - "device": device_hint, - }, - ) - return LoadedUltralyticsDetector( - detector_model=detector_model, - bbox_detector=bbox_detector, - segm_detector=segm_detector, - ) - - def reject_sentinel(self, model_name: str) -> None: - """Reject placeholder dropdown selections before filesystem work.""" - - if model_name == NO_LOCAL_ULTRALYTICS_MODELS: - raise ValueError( - "No local Ultralytics models are available. Enable 'Show " - "downloadable models in loader dropdowns' in SimpleSyrup settings " - "or install a model in models\\ultralytics, " - "models\\ultralytics\\bbox, or models\\ultralytics\\segm." - ) - - def resolve_model_path(self, model_name: str) -> Path: - """Resolve a safe model choice to a file inside ComfyUI model folders.""" - - safe_name = Path(model_name.replace("\\", "/")) - if safe_name.is_absolute() or ".." in safe_name.parts: - raise ValueError( - f"Ultralytics model name '{model_name}' is not a safe relative path." - ) - - candidates = self._candidate_paths(safe_name) - for candidate in candidates: - if candidate.is_file(): - return candidate - - raise ValueError( - f"Ultralytics model '{model_name}' was not found in configured " - "ComfyUI model folders." - ) - - def _candidate_paths(self, model_name: Path) -> list[Path]: - """Return bounded filesystem candidates for a model choice.""" - - self._register_model_folders() - candidates: list[Path] = [] - parts = model_name.parts - if len(parts) >= 2 and parts[0] == "bbox": - relative = Path(*parts[1:]) - candidates.extend( - folder / relative - for folder in self._folder_paths_for(ULTRALYTICS_BBOX_FOLDER) - ) - candidates.extend( - folder / "bbox" / relative - for folder in self._folder_paths_for(ULTRALYTICS_FOLDER) - ) - elif len(parts) >= 2 and parts[0] == "segm": - relative = Path(*parts[1:]) - candidates.extend( - folder / relative - for folder in self._folder_paths_for(ULTRALYTICS_SEGM_FOLDER) - ) - candidates.extend( - folder / "segm" / relative - for folder in self._folder_paths_for(ULTRALYTICS_FOLDER) - ) - else: - candidates.extend( - folder / model_name - for folder in self._folder_paths_for(ULTRALYTICS_FOLDER) - ) - return candidates - - def _register_model_folders(self) -> None: - """Register conventional Ultralytics folders with ComfyUI when possible.""" - - folder_paths = self._folder_paths() - models_dir = Path(str(folder_paths.models_dir)) - add_model_folder_path = getattr(folder_paths, "add_model_folder_path", None) - if add_model_folder_path is None: - return - - registry = cast( - ModelFolderRegistry, - getattr(folder_paths, "folder_names_and_paths", {}), - ) - registrations = ( - (ULTRALYTICS_FOLDER, models_dir / "ultralytics"), - (ULTRALYTICS_BBOX_FOLDER, models_dir / "ultralytics" / "bbox"), - (ULTRALYTICS_SEGM_FOLDER, models_dir / "ultralytics" / "segm"), - ) - for folder_name, path in registrations: - if folder_name in registry: - continue - add_model_folder_path(folder_name, str(path)) - - def _folder_paths_for(self, folder_name: str) -> list[Path]: - """Return registered paths for one ComfyUI model folder.""" - - folder_paths = self._folder_paths() - models_dir = Path(str(folder_paths.models_dir)) - fallback = { - ULTRALYTICS_FOLDER: models_dir / "ultralytics", - ULTRALYTICS_BBOX_FOLDER: models_dir / "ultralytics" / "bbox", - ULTRALYTICS_SEGM_FOLDER: models_dir / "ultralytics" / "segm", - }[folder_name] - registry = cast( - ModelFolderRegistry, - getattr(folder_paths, "folder_names_and_paths", {}), - ) - paths = [fallback] - if folder_name in registry: - paths = [Path(str(path)) for path in registry[folder_name][0]] + paths - return _unique_paths(paths) - - def _folder_paths(self) -> ModuleType: - """Import ComfyUI folder path helpers lazily.""" - - if self._folder_paths_module is not None: - return self._folder_paths_module - module = importlib.import_module("folder_paths") - if not isinstance(module, ModuleType): - raise TypeError("folder_paths import did not return a module.") - self._folder_paths_module = module - return module - - def _ultralytics(self) -> ModuleType: - """Import Ultralytics lazily and fail with an actionable message.""" - - if self._ultralytics_module is not None: - return self._ultralytics_module - try: - module = importlib.import_module("ultralytics") - except ModuleNotFoundError as exc: - raise RuntimeError( - "Ultralytics support requires the 'ultralytics' package in the " - "ComfyUI virtual environment." - ) from exc - if not isinstance(module, ModuleType): - raise TypeError("ultralytics import did not return a module.") - self._ultralytics_module = module - return module - - -def _supported_files(folder: Path) -> list[Path]: - """Return directly contained supported model files for a folder.""" - - if not folder.is_dir(): - return [] - return sorted( - path - for path in folder.iterdir() - if path.is_file() and path.suffix.lower() in SUPPORTED_MODEL_EXTENSIONS - ) - - -def _unique_paths(paths: list[Path]) -> list[Path]: - """Return unique paths while preserving order.""" - - unique: list[Path] = [] - seen: set[str] = set() - for path in paths: - key = str(path) - if key in seen: - continue - unique.append(path) - seen.add(key) - return unique - - -def _normalized_model_name(model_name: str) -> str: - """Return a stable model selection string for cache identity.""" - - return model_name.replace("\\", "/") - - -def _catalog_entry_or_none(selection: str) -> ModelEntry | None: - """Return a curated Ultralytics entry when a dropdown label matches it.""" - - return next( - ( - entry - for entry in ULTRALYTICS_ENTRIES - if selection in (entry.entry_id, entry.display_name) - ), - None, - ) - - -def _catalog_selection(entry: ModelEntry) -> str: - """Return the local conventional selection path for one catalog entry.""" - - if len(entry.artifacts) != 1: - raise ValueError( - f"Ultralytics catalog entry '{entry.entry_id}' must have one artifact." - ) - - artifact = entry.artifacts[0] - prefix_by_folder = { - ULTRALYTICS_BBOX_FOLDER: "bbox", - ULTRALYTICS_SEGM_FOLDER: "segm", - } - try: - prefix = prefix_by_folder[artifact.folder_name] - except KeyError as error: - raise ValueError( - f"Ultralytics catalog entry '{entry.entry_id}' has unsupported folder " - f"'{artifact.folder_name}'." - ) from error - return f"{prefix}/{artifact.filename}" - - -def _model_task(model_name: str, raw_model: object) -> str: - """Infer detector task from choice prefix or model metadata.""" - - normalized_name = model_name.replace("\\", "/") - if normalized_name.startswith("segm/"): - return "segment" - if normalized_name.startswith("bbox/"): - return "detect" - - task = getattr(raw_model, "task", None) - if isinstance(task, str) and task: - return task - return "detect" - - -def _model_names(raw_model: object) -> dict[int, str]: - """Extract class names from a loaded Ultralytics model.""" - - names = getattr(raw_model, "names", {}) - if isinstance(names, dict): - return {int(key): str(value) for key, value in names.items()} - if isinstance(names, list): - return {index: str(value) for index, value in enumerate(names)} - return {} - - -def _model_device_hint(raw_model: object) -> str: - """Return a best-effort Ultralytics device hint for diagnostics.""" - - direct_device = getattr(raw_model, "device", None) - if direct_device is not None: - return str(direct_device) - inner_model = getattr(raw_model, "model", None) - inner_device = getattr(inner_model, "device", None) - if inner_device is not None: - return str(inner_device) - return "runtime-owned" diff --git a/simple_syrup/runtime/ultralytics_model_adapter.py b/simple_syrup/runtime/ultralytics_model_adapter.py new file mode 100644 index 0000000..09e6972 --- /dev/null +++ b/simple_syrup/runtime/ultralytics_model_adapter.py @@ -0,0 +1,141 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Load Ultralytics models behind a narrow runtime adapter.""" + +from __future__ import annotations + +import importlib +from dataclasses import dataclass +from pathlib import Path +from types import ModuleType +from typing import Any + +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) + + +@dataclass(frozen=True) +class UltralyticsDetectorModel: + """Store a loaded Ultralytics detector with SimpleSyrup metadata.""" + + model_name: str + model_path: Path + model: Any + task: str + names: dict[int, str] + supports_segmentation: bool + + +class UltralyticsModelAdapter: + """Construct detector models through the optional Ultralytics runtime.""" + + def __init__(self, ultralytics_module: ModuleType | None = None) -> None: + """Create the adapter with an optional runtime module override.""" + + self._ultralytics_module = ultralytics_module + + def load(self, model_name: str, model_path: Path) -> UltralyticsDetectorModel: + """Load one detector checkpoint and expose normalized metadata.""" + + model_class = getattr(self._ultralytics(), "YOLO", None) + if model_class is None: + raise RuntimeError( + "Ultralytics support requires a module exposing the YOLO class." + ) + + try: + raw_model = model_class(str(model_path)) + except Exception as exc: + LOGGER.error( + "Failed to load Ultralytics model", + extra={ + "operation": "load_ultralytics_model", + "model_name": model_name, + "model_path": str(model_path), + }, + exc_info=True, + ) + raise RuntimeError( + f"Ultralytics model '{model_name}' could not be loaded from " + f"'{model_path}'." + ) from exc + + task = _model_task(model_name, raw_model) + detector_model = UltralyticsDetectorModel( + model_name=model_name, + model_path=model_path, + model=raw_model, + task=task, + names=_model_names(raw_model), + supports_segmentation=task in {"segment", "segm"}, + ) + LOGGER.info( + "Ultralytics model loaded", + extra={ + "operation": "load_ultralytics_model", + "model_name": model_name, + "model_path": str(model_path), + "task": task, + "device": _model_device_hint(raw_model), + }, + ) + return detector_model + + def _ultralytics(self) -> ModuleType: + """Import Ultralytics lazily and fail with an actionable message.""" + + if self._ultralytics_module is not None: + return self._ultralytics_module + try: + module = importlib.import_module("ultralytics") + except ModuleNotFoundError as exc: + raise RuntimeError( + "Ultralytics support requires the 'ultralytics' package in the " + "ComfyUI virtual environment." + ) from exc + if not isinstance(module, ModuleType): + raise TypeError("ultralytics import did not return a module.") + self._ultralytics_module = module + return module + + +def _model_task(model_name: str, raw_model: object) -> str: + """Infer detector task from choice prefix or model metadata.""" + + normalized_name = model_name.replace("\\", "/") + if normalized_name.startswith("segm/"): + return "segment" + if normalized_name.startswith("bbox/"): + return "detect" + + task = getattr(raw_model, "task", None) + if isinstance(task, str) and task: + return task + return "detect" + + +def _model_names(raw_model: object) -> dict[int, str]: + """Extract class names from a loaded Ultralytics model.""" + + names = getattr(raw_model, "names", {}) + if isinstance(names, dict): + return {int(key): str(value) for key, value in names.items()} + if isinstance(names, list): + return {index: str(value) for index, value in enumerate(names)} + return {} + + +def _model_device_hint(raw_model: object) -> str: + """Return a best-effort Ultralytics device hint for diagnostics.""" + + direct_device = getattr(raw_model, "device", None) + if direct_device is not None: + return str(direct_device) + inner_model = getattr(raw_model, "model", None) + inner_device = getattr(inner_model, "device", None) + if inner_device is not None: + return str(inner_device) + return "runtime-owned" diff --git a/simple_syrup/runtime/ultralytics_model_folders.py b/simple_syrup/runtime/ultralytics_model_folders.py new file mode 100644 index 0000000..4d88eee --- /dev/null +++ b/simple_syrup/runtime/ultralytics_model_folders.py @@ -0,0 +1,203 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover and resolve Ultralytics checkpoints in ComfyUI model folders.""" + +from __future__ import annotations + +import importlib +from pathlib import Path +from types import ModuleType +from typing import TypeAlias, cast + +from .model_folders import SUPPORTED_MODEL_EXTENSIONS + +ULTRALYTICS_FOLDER = "ultralytics" +ULTRALYTICS_BBOX_FOLDER = "ultralytics_bbox" +ULTRALYTICS_SEGM_FOLDER = "ultralytics_segm" + +ModelFolderRegistry: TypeAlias = dict[str, tuple[list[str], set[str]]] + + +class UltralyticsModelFolders: + """Own ComfyUI folder registration, discovery, and safe path resolution.""" + + def __init__(self, folder_paths_module: ModuleType | None = None) -> None: + """Create the adapter with an optional ComfyUI module override.""" + + self._folder_paths_module = folder_paths_module + + @property + def folder_paths_module(self) -> ModuleType | None: + """Return the resolved or injected folder-paths module when available.""" + + return self._folder_paths_module + + def available_models(self) -> list[str]: + """Return supported model files in registered Ultralytics folders.""" + + self.register() + folder_paths = self._folder_paths() + choices: set[str] = set() + + for folder in self.paths_for(ULTRALYTICS_FOLDER): + if not folder.is_dir(): + continue + choices.update(path.name for path in _supported_files(folder)) + choices.update( + f"bbox/{path.name}" for path in _supported_files(folder / "bbox") + ) + choices.update( + f"segm/{path.name}" for path in _supported_files(folder / "segm") + ) + + for path in self.paths_for(ULTRALYTICS_BBOX_FOLDER): + choices.update(f"bbox/{file.name}" for file in _supported_files(path)) + for path in self.paths_for(ULTRALYTICS_SEGM_FOLDER): + choices.update(f"segm/{file.name}" for file in _supported_files(path)) + + registry = cast( + ModelFolderRegistry, + getattr(folder_paths, "folder_names_and_paths", {}), + ) + for folder_name in ( + ULTRALYTICS_FOLDER, + ULTRALYTICS_BBOX_FOLDER, + ULTRALYTICS_SEGM_FOLDER, + ): + if folder_name not in registry: + continue + for filename in folder_paths.get_filename_list(folder_name): + path = Path(str(filename)) + if path.suffix.lower() not in SUPPORTED_MODEL_EXTENSIONS: + continue + if folder_name == ULTRALYTICS_BBOX_FOLDER: + choices.add(f"bbox/{path.name}") + elif folder_name == ULTRALYTICS_SEGM_FOLDER: + choices.add(f"segm/{path.name}") + else: + choices.add(path.as_posix()) + return sorted(choices) + + def resolve(self, model_name: str) -> Path: + """Resolve a safe model choice inside the configured model folders.""" + + safe_name = Path(model_name.replace("\\", "/")) + if safe_name.is_absolute() or ".." in safe_name.parts: + raise ValueError( + f"Ultralytics model name '{model_name}' is not a safe relative path." + ) + for candidate in self._candidate_paths(safe_name): + if candidate.is_file(): + return candidate + raise ValueError( + f"Ultralytics model '{model_name}' was not found in configured " + "ComfyUI model folders." + ) + + def register(self) -> None: + """Register conventional Ultralytics folders with ComfyUI when possible.""" + + folder_paths = self._folder_paths() + models_dir = Path(str(folder_paths.models_dir)) + add_model_folder_path = getattr(folder_paths, "add_model_folder_path", None) + if add_model_folder_path is None: + return + registry = cast( + ModelFolderRegistry, + getattr(folder_paths, "folder_names_and_paths", {}), + ) + registrations = ( + (ULTRALYTICS_FOLDER, models_dir / "ultralytics"), + (ULTRALYTICS_BBOX_FOLDER, models_dir / "ultralytics" / "bbox"), + (ULTRALYTICS_SEGM_FOLDER, models_dir / "ultralytics" / "segm"), + ) + for folder_name, path in registrations: + if folder_name not in registry: + add_model_folder_path(folder_name, str(path)) + + def paths_for(self, folder_name: str) -> list[Path]: + """Return registered paths for one ComfyUI model folder.""" + + folder_paths = self._folder_paths() + models_dir = Path(str(folder_paths.models_dir)) + fallback = { + ULTRALYTICS_FOLDER: models_dir / "ultralytics", + ULTRALYTICS_BBOX_FOLDER: models_dir / "ultralytics" / "bbox", + ULTRALYTICS_SEGM_FOLDER: models_dir / "ultralytics" / "segm", + }[folder_name] + registry = cast( + ModelFolderRegistry, + getattr(folder_paths, "folder_names_and_paths", {}), + ) + paths = [fallback] + if folder_name in registry: + paths = [Path(str(path)) for path in registry[folder_name][0]] + paths + return _unique_paths(paths) + + def _candidate_paths(self, model_name: Path) -> list[Path]: + """Return bounded filesystem candidates for a model choice.""" + + self.register() + parts = model_name.parts + if len(parts) >= 2 and parts[0] == "bbox": + relative = Path(*parts[1:]) + return [ + *( + folder / relative + for folder in self.paths_for(ULTRALYTICS_BBOX_FOLDER) + ), + *( + folder / "bbox" / relative + for folder in self.paths_for(ULTRALYTICS_FOLDER) + ), + ] + if len(parts) >= 2 and parts[0] == "segm": + relative = Path(*parts[1:]) + return [ + *( + folder / relative + for folder in self.paths_for(ULTRALYTICS_SEGM_FOLDER) + ), + *( + folder / "segm" / relative + for folder in self.paths_for(ULTRALYTICS_FOLDER) + ), + ] + return [folder / model_name for folder in self.paths_for(ULTRALYTICS_FOLDER)] + + def _folder_paths(self) -> ModuleType: + """Import ComfyUI folder path helpers lazily.""" + + if self._folder_paths_module is None: + module = importlib.import_module("folder_paths") + if not isinstance(module, ModuleType): + raise TypeError("folder_paths import did not return a module.") + self._folder_paths_module = module + return self._folder_paths_module + + +def _supported_files(folder: Path) -> list[Path]: + """Return directly contained supported model files for a folder.""" + + if not folder.is_dir(): + return [] + return sorted( + path + for path in folder.iterdir() + if path.is_file() and path.suffix.lower() in SUPPORTED_MODEL_EXTENSIONS + ) + + +def _unique_paths(paths: list[Path]) -> list[Path]: + """Return unique paths while preserving order.""" + + unique: list[Path] = [] + seen: set[str] = set() + for path in paths: + key = str(path) + if key not in seen: + unique.append(path) + seen.add(key) + return unique diff --git a/simple_syrup/services/attention_coupling_preparation_service.py b/simple_syrup/services/attention_coupling_preparation_service.py index 56562bd..3bbc020 100644 --- a/simple_syrup/services/attention_coupling_preparation_service.py +++ b/simple_syrup/services/attention_coupling_preparation_service.py @@ -11,6 +11,7 @@ from dataclasses import dataclass import torch +from ..domain import attention_coupling_preparation as preparation_domain from ..domain.raw_regional_attention import ( RawRegionalAttentionBranch, RawRegionalAttentionPlan, @@ -31,15 +32,6 @@ _UNSUPPORTED_METADATA_KEYS = frozenset( ) -@dataclass(frozen=True, slots=True) -class AttentionCouplingPreparation: - """Retain the full raw plan and base-only ordinary sampler inputs.""" - - plan: RawRegionalAttentionPlan - positive: object - negative: object - - @dataclass(frozen=True, slots=True) class AttentionCouplingMetadataIssue: """Identify one unsupported or malformed conditioning metadata field.""" @@ -71,7 +63,9 @@ class AttentionCouplingMetadataError(ValueError): class AttentionCouplingPreparationService: """Separate regional contexts from ordinary KSampler conditioning.""" - def prepare(self, plan: RawRegionalAttentionPlan) -> AttentionCouplingPreparation: + def prepare( + self, plan: RawRegionalAttentionPlan + ) -> preparation_domain.AttentionCouplingPreparation: """Validate all contexts before returning the untouched base conditionings.""" if not isinstance(plan, RawRegionalAttentionPlan): @@ -82,7 +76,7 @@ class AttentionCouplingPreparationService: ) if issues: raise AttentionCouplingMetadataError(issues) - return AttentionCouplingPreparation( + return preparation_domain.AttentionCouplingPreparation( plan=plan, positive=plan.positive.base_conditioning, negative=plan.negative.base_conditioning, diff --git a/simple_syrup/services/attention_region_matte.py b/simple_syrup/services/attention_region_matte.py index e636f76..b51be27 100644 --- a/simple_syrup/services/attention_region_matte.py +++ b/simple_syrup/services/attention_region_matte.py @@ -8,8 +8,8 @@ from __future__ import annotations import torch +from ..domain.segs_mask_ops import feather_mask from ..masking.mask_components import connected_mask_components -from ..masking.segs_mask_ops import feather_mask class AttentionMatteService: diff --git a/simple_syrup/services/detail_segs_as_regions_service.py b/simple_syrup/services/detail_segs_as_regions_service.py index 9276b0b..f178adb 100644 --- a/simple_syrup/services/detail_segs_as_regions_service.py +++ b/simple_syrup/services/detail_segs_as_regions_service.py @@ -21,6 +21,7 @@ from ..domain.regional_detailing import ( pair_segments_with_conditioning, ) from ..domain.segs import CropRegion, coerce_segs +from ..domain.segs_mask_ops import validate_single_image from ..masking.regional_detailing_masks import ( build_image_regions, build_latent_regions, @@ -28,11 +29,10 @@ from ..masking.regional_detailing_masks import ( scale_image_regions, union_masks, ) -from ..masking.segs_mask_ops import validate_single_image -from ..runtime import regional_multidiffusion_sampling from ..runtime.detail_previews import DetailPreviewContext, work_region_from_mask from ..runtime.detail_resize import DetailImageResizer -from ..runtime.detail_sampling import DetailSampler, Latent +from ..runtime.detail_sampling import Latent +from ..runtime.regional_detail_sampler import RegionalDetailSampler from ..shared.logging import get_logger LOGGER = get_logger(__name__) @@ -97,62 +97,6 @@ class DetailSEGSAsRegionsResult: image: torch.Tensor -class RegionalDetailSampler: - """Adapt shared detail sampling helpers to regional MultiDiffusion.""" - - def __init__(self, detail_sampler: DetailSampler | None = None) -> None: - """Create the runtime adapter with injectable encode/decode helper.""" - - self._detail_sampler = detail_sampler or DetailSampler() - - def encode(self, vae: Any, pixels: torch.Tensor, tiled: bool) -> Latent: - """Encode pixels into a latent dictionary.""" - - return self._detail_sampler.encode(vae, pixels, tiled) - - def decode(self, vae: Any, latent: Latent, tiled: bool) -> torch.Tensor: - """Decode latent samples into pixels.""" - - return self._detail_sampler.decode(vae, latent, tiled) - - def sample_regions( - self, - *, - model: Any, - seed: int, - steps: int, - cfg: float, - sampler_name: str, - scheduler: str, - positive: Any, - negative: Any, - latent_image: Latent, - regions: tuple[LatentRegion, ...], - denoise: float, - global_prompt_weight: float, - preview_context: DetailPreviewContext | None = None, - differential_diffusion: bool = False, - ) -> Latent: - """Sample one full latent with regional MultiDiffusion.""" - - return regional_multidiffusion_sampling.sample_regional_multidiffusion( - model=model, - seed=seed, - steps=steps, - cfg=cfg, - sampler_name=sampler_name, - scheduler=scheduler, - positive=positive, - negative=negative, - latent_image=latent_image, - regions=regions, - denoise=denoise, - global_prompt_weight=global_prompt_weight, - preview_context=preview_context, - differential_diffusion=differential_diffusion, - ) - - class DetailSEGSAsRegionsService: """Detail provided SEGS through one regional MultiDiffusion pass.""" diff --git a/simple_syrup/services/detail_segs_by_scale_factor_service.py b/simple_syrup/services/detail_segs_by_scale_factor_service.py index b4cf9de..4ef3d61 100644 --- a/simple_syrup/services/detail_segs_by_scale_factor_service.py +++ b/simple_syrup/services/detail_segs_by_scale_factor_service.py @@ -14,12 +14,12 @@ import torch from ..domain.conditioning_batch import select_conditioning from ..domain.detail_geometry import DetailScalePlan, build_detail_scale_plan from ..domain.segs import Segment, coerce_segs -from ..image.crop_composite import composite_crop -from ..masking.detailer_masks import gaussian_feather_mask -from ..masking.segs_mask_ops import ( +from ..domain.segs_mask_ops import ( crop_image, validate_single_image, ) +from ..image.crop_composite import composite_crop +from ..masking.detailer_masks import gaussian_feather_mask from ..runtime.detail_previews import DetailPreviewContext from ..runtime.detail_resize import DetailImageResizer from ..runtime.detail_sampling import DetailSampler, Latent diff --git a/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py b/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py index 1670c22..9765d29 100644 --- a/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py +++ b/simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py @@ -14,49 +14,22 @@ import torch from ..domain.conditioning_batch import select_conditioning from ..domain.detail_geometry import DetailScalePlan, build_detail_scale_plan from ..domain.segs import Segment, coerce_segs -from ..domain.tiled_diffusion import validate_tiled_diffusion_mode -from ..image.crop_composite import composite_crop -from ..masking.detailer_masks import gaussian_feather_mask -from ..masking.segs_mask_ops import ( +from ..domain.segs_mask_ops import ( crop_image, validate_single_image, ) +from ..domain.tiled_diffusion import validate_tiled_diffusion_mode +from ..image.crop_composite import composite_crop +from ..masking.detailer_masks import gaussian_feather_mask from ..runtime.detail_previews import DetailPreviewContext from ..runtime.detail_resize import DetailImageResizer -from ..runtime.detail_sampling import DetailSampler, Latent +from ..runtime.detail_sampling import Latent from ..shared.logging import get_logger -from .tiled_diffusion_sampling_service import TiledDiffusionSamplingService +from .tiled_detail_sampler import TiledDetailSampler LOGGER = get_logger(__name__) -class TiledDiffusionLatentSamplingBoundary(Protocol): - """Latent sampling boundary for selectable tiled diffusion modes.""" - - def sample( - 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 = None, - differential_diffusion: bool = False, - ) -> Latent: - """Sample a latent using the selected tiled diffusion mode.""" - - class TiledDetailSamplingBoundary(Protocol): """Runtime boundary used by tiled scale-factor detailing.""" @@ -118,75 +91,6 @@ class TiledDetailerResult: image: torch.Tensor -class TiledDetailSampler: - """Adapt shared VAE helpers and tiled diffusion runtimes.""" - - def __init__( - self, - detail_sampler: DetailSampler | None = None, - tiled_sampling_service: TiledDiffusionLatentSamplingBoundary | None = None, - ) -> None: - """Create the adapter with injectable encode/decode behavior.""" - - self._detail_sampler = detail_sampler or DetailSampler() - self._tiled_sampling_service = ( - tiled_sampling_service or TiledDiffusionSamplingService() - ) - - def encode(self, vae: Any, pixels: torch.Tensor, tiled: bool) -> Latent: - """Encode pixels into a latent dictionary.""" - - return self._detail_sampler.encode(vae, pixels, tiled) - - def decode(self, vae: Any, latent: Latent, tiled: bool) -> torch.Tensor: - """Decode latent samples into pixels.""" - - return self._detail_sampler.decode(vae, latent, tiled) - - def sample_tiled( - 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 = None, - differential_diffusion: bool = False, - ) -> Latent: - """Sample one latent crop with the selected tiled diffusion runtime.""" - - return self._tiled_sampling_service.sample( - 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, - ) - - class DetailSEGSByScaleFactorTiledDiffusionService: """Detail SEGS crops with tiled diffusion latent sampling.""" diff --git a/simple_syrup/runtime/detector_compat.py b/simple_syrup/services/detector_compat.py similarity index 75% rename from simple_syrup/runtime/detector_compat.py rename to simple_syrup/services/detector_compat.py index 3a9fd79..c64d4cd 100644 --- a/simple_syrup/runtime/detector_compat.py +++ b/simple_syrup/services/detector_compat.py @@ -2,14 +2,15 @@ # Copyright (C) 2026 Artificial Sweetener and contributors # SPDX-License-Identifier: AGPL-3.0-or-later -"""Impact-style detector facades backed by native SimpleSyrup services.""" +"""Expose Impact-compatible detector facades over native detection services.""" from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from ..domain.segs import to_impact_compatible_segs -from .ultralytics_loader import UltralyticsDetectorModel +from ..runtime.ultralytics_model_adapter import UltralyticsDetectorModel +from .segs_detection_service import SegsDetectionService @dataclass(frozen=True) @@ -17,6 +18,11 @@ class BBoxDetectorFacade: """Expose a bbox detector-shaped object for existing workflows.""" detector_model: UltralyticsDetectorModel + detection_service: SegsDetectionService = field( + default_factory=SegsDetectionService, + repr=False, + compare=False, + ) def detect( self, @@ -30,9 +36,7 @@ class BBoxDetectorFacade: """Detect rectangular SEGS through the native detection service.""" del detailer_hook - from ..services.segs_detection_service import SegsDetectionService - - segs = SegsDetectionService().detect( + segs = self.detection_service.detect( image=image, detector_model=self.detector_model, threshold=threshold, @@ -50,6 +54,11 @@ class SegmDetectorFacade: detector_model: UltralyticsDetectorModel bbox_detector: BBoxDetectorFacade + detection_service: SegsDetectionService = field( + default_factory=SegsDetectionService, + repr=False, + compare=False, + ) def detect( self, @@ -63,9 +72,7 @@ class SegmDetectorFacade: """Detect segmentation SEGS when available, otherwise rectangular SEGS.""" del detailer_hook - from ..services.segs_detection_service import SegsDetectionService - - segs = SegsDetectionService().detect( + segs = self.detection_service.detect( image=image, detector_model=self.detector_model, threshold=threshold, diff --git a/simple_syrup/services/mask_to_segs_service.py b/simple_syrup/services/mask_to_segs_service.py index cda0e4e..861dcc6 100644 --- a/simple_syrup/services/mask_to_segs_service.py +++ b/simple_syrup/services/mask_to_segs_service.py @@ -9,14 +9,14 @@ from __future__ import annotations import torch from ..domain.segs import NativeSegs, Segment -from ..masking.mask_components import connected_mask_components -from ..masking.segs_mask_ops import ( +from ..domain.segs_mask_ops import ( crop_image, crop_mask, crop_region_for_bbox, dilate_mask, validate_single_image, ) +from ..masking.mask_components import connected_mask_components from ..shared.logging import get_logger LOGGER = get_logger(__name__) diff --git a/simple_syrup/runtime/prompt_control_batch_graph.py b/simple_syrup/services/prompt_control_batch_graph.py similarity index 97% rename from simple_syrup/runtime/prompt_control_batch_graph.py rename to simple_syrup/services/prompt_control_batch_graph.py index 67fd9b7..7e9cee8 100644 --- a/simple_syrup/runtime/prompt_control_batch_graph.py +++ b/simple_syrup/services/prompt_control_batch_graph.py @@ -9,13 +9,13 @@ from __future__ import annotations from typing import Any from ..domain.prompt_control_prompt import PreparedPromptSide -from ..services.prompt_control_segment_planning_service import ( - PromptControlSegmentPlanningService, -) -from .prompt_control_graph_adapter import ( +from ..runtime.prompt_control_graph_adapter import ( PromptControlGraphAdapter, RegionalSegmentEncoding, ) +from .prompt_control_segment_planning_service import ( + PromptControlSegmentPlanningService, +) PROMPT_CONTROL_MISSING_MESSAGE = ( "Encode Prompt Batch w/ Prompt Control requires comfyui-prompt-control. " diff --git a/simple_syrup/runtime/prompt_control_schedule_encode_graph.py b/simple_syrup/services/prompt_control_schedule_encode_graph.py similarity index 98% rename from simple_syrup/runtime/prompt_control_schedule_encode_graph.py rename to simple_syrup/services/prompt_control_schedule_encode_graph.py index 0a8cd71..396c7a8 100644 --- a/simple_syrup/runtime/prompt_control_schedule_encode_graph.py +++ b/simple_syrup/services/prompt_control_schedule_encode_graph.py @@ -11,14 +11,14 @@ from typing import Any from ..domain.negative_prompt_weights import contains_negative_prompt_weight from ..domain.prompt_batch_parser import DEFAULT_PROMPT_BATCH_SEPARATOR from ..domain.prompt_control_prompt import PreparedPromptSide, apply_encode_style -from ..services.prompt_control_segment_planning_service import ( - PromptControlSegmentPlan, - PromptControlSegmentPlanningService, -) -from .prompt_control_graph_adapter import ( +from ..runtime.prompt_control_graph_adapter import ( PromptControlGraphAdapter, RegionalSegmentEncoding, ) +from .prompt_control_segment_planning_service import ( + PromptControlSegmentPlan, + PromptControlSegmentPlanningService, +) PROMPT_CONTROL_MISSING_MESSAGE = ( "Schedule & Encode Prompts requires comfyui-prompt-control. " diff --git a/simple_syrup/services/segs_detection_service.py b/simple_syrup/services/segs_detection_service.py index bdb33af..8ebd076 100644 --- a/simple_syrup/services/segs_detection_service.py +++ b/simple_syrup/services/segs_detection_service.py @@ -11,7 +11,7 @@ from typing import Protocol import torch from ..domain.segs import NativeSegs, Segment, coerce_segment_mask -from ..masking.segs_mask_ops import ( +from ..domain.segs_mask_ops import ( crop_image, crop_mask, crop_region_for_bbox, @@ -23,7 +23,7 @@ from ..runtime.ultralytics_detection import ( UltralyticsDetection, run_ultralytics_detection, ) -from ..runtime.ultralytics_loader import UltralyticsDetectorModel +from ..runtime.ultralytics_model_adapter import UltralyticsDetectorModel from ..shared.logging import get_logger from .segs_output_service import combined_mask_from_segs diff --git a/simple_syrup/services/segs_from_sam_output_service.py b/simple_syrup/services/segs_from_sam_output_service.py index 9a7ed92..24601c6 100644 --- a/simple_syrup/services/segs_from_sam_output_service.py +++ b/simple_syrup/services/segs_from_sam_output_service.py @@ -14,7 +14,7 @@ import torch import torch.nn.functional as functional from ..domain.segs import BoundingBox, CropRegion, NativeSegs, Segment -from ..masking.segs_mask_ops import validate_single_image +from ..domain.segs_mask_ops import validate_single_image from ..runtime.progress import NullPhaseProgressReporter, PhaseProgressReporter from ..runtime.sam_automatic_segmenter import ( AutomaticSAMMask, diff --git a/simple_syrup/services/segs_output_service.py b/simple_syrup/services/segs_output_service.py index bd2d432..3084dc0 100644 --- a/simple_syrup/services/segs_output_service.py +++ b/simple_syrup/services/segs_output_service.py @@ -21,7 +21,7 @@ from ..domain.segs import ( sort_segs, to_impact_compatible_segs, ) -from ..masking.segs_mask_ops import ( +from ..domain.segs_mask_ops import ( crop_image, crop_mask, crop_region_for_bbox, diff --git a/simple_syrup/services/simple_preview_segs_service.py b/simple_syrup/services/simple_preview_segs_service.py index dd47cc4..b391f44 100644 --- a/simple_syrup/services/simple_preview_segs_service.py +++ b/simple_syrup/services/simple_preview_segs_service.py @@ -6,18 +6,18 @@ from __future__ import annotations -from dataclasses import dataclass from math import ceil, sqrt import torch import torch.nn.functional as functional +from ..domain.seg_preview import AtlasPlacement, SegPreviewDocument, SegPreviewRegion from ..domain.seg_visualization import ( SegVisualizationPlan, build_seg_visualization_plan, ) -from ..domain.segs import CropRegion, NativeSegs -from ..masking.segs_mask_ops import validate_single_image +from ..domain.segs import NativeSegs +from ..domain.segs_mask_ops import validate_single_image _MAX_PREVIEW_EDGE = 1024 _ATLAS_PIXEL_BUDGET = 4 * 1024 * 1024 @@ -25,44 +25,6 @@ _MAX_ATLAS_EDGE = 2048 _ATLAS_PADDING = 1 -@dataclass(frozen=True) -class AtlasPlacement: - """Locate one region mask inside the packed mask atlas.""" - - left: int - top: int - width: int - height: int - - -@dataclass(frozen=True) -class SegPreviewRegion: - """Describe one interactive region and its packed mask geometry.""" - - region_id: str - index: int - label: str - confidence: float - active_area: int - color: str - crop: CropRegion - atlas: AtlasPlacement - - -@dataclass(frozen=True) -class SegPreviewDocument: - """Carry bounded image assets and interaction metadata to the UI adapter.""" - - source_width: int - source_height: int - preview_width: int - preview_height: int - image: torch.Tensor - atlas: torch.Tensor - region_images: tuple[torch.Tensor, ...] - regions: tuple[SegPreviewRegion, ...] - - class SimplePreviewSEGSService: """Create one compact interactive-preview document without ComfyUI IO.""" diff --git a/simple_syrup/services/tag_segs_with_external_llm_service.py b/simple_syrup/services/tag_segs_with_external_llm_service.py index 2db45df..90dd3cb 100644 --- a/simple_syrup/services/tag_segs_with_external_llm_service.py +++ b/simple_syrup/services/tag_segs_with_external_llm_service.py @@ -25,7 +25,7 @@ from ..domain.segs import ( coerce_segs, to_impact_compatible_segs, ) -from ..masking.segs_mask_ops import validate_single_image +from ..domain.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 diff --git a/simple_syrup/services/tag_segs_with_wd14_service.py b/simple_syrup/services/tag_segs_with_wd14_service.py index a522fa0..a604553 100644 --- a/simple_syrup/services/tag_segs_with_wd14_service.py +++ b/simple_syrup/services/tag_segs_with_wd14_service.py @@ -21,7 +21,7 @@ from ..domain.segs import ( coerce_segs, to_impact_compatible_segs, ) -from ..masking.segs_mask_ops import crop_image, validate_single_image +from ..domain.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 diff --git a/simple_syrup/services/tile_and_tag_segs_service.py b/simple_syrup/services/tile_and_tag_segs_service.py index 80543c0..e430c21 100644 --- a/simple_syrup/services/tile_and_tag_segs_service.py +++ b/simple_syrup/services/tile_and_tag_segs_service.py @@ -15,8 +15,8 @@ import torch from ..domain.conditioning_batch import ConditioningBatch from ..domain.prompt_composition import prefix_prompt from ..domain.segs import ImpactSegs, NativeSegs, to_impact_compatible_segs +from ..domain.segs_mask_ops import crop_image, validate_single_image from ..domain.tile_segs import TileSEGSBuilder, TileSEGSControls -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 diff --git a/simple_syrup/services/tiled_detail_sampler.py b/simple_syrup/services/tiled_detail_sampler.py new file mode 100644 index 0000000..19a346f --- /dev/null +++ b/simple_syrup/services/tiled_detail_sampler.py @@ -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 + +"""Adapt detail encode/decode behavior to tiled diffusion sampling.""" + +from __future__ import annotations + +from typing import Any, Protocol + +import torch + +from ..runtime.detail_previews import DetailPreviewContext +from ..runtime.detail_sampling import DetailSampler, Latent +from .tiled_diffusion_sampling_service import TiledDiffusionSamplingService + + +class TiledDiffusionLatentSamplingBoundary(Protocol): + """Latent sampling boundary for selectable tiled diffusion modes.""" + + def sample( + 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 = None, + differential_diffusion: bool = False, + ) -> Latent: + """Sample a latent using the selected tiled diffusion mode.""" + + +class TiledDetailSampler: + """Adapt shared VAE helpers and tiled diffusion application services.""" + + def __init__( + self, + detail_sampler: DetailSampler | None = None, + tiled_sampling_service: TiledDiffusionLatentSamplingBoundary | None = None, + ) -> None: + """Create the adapter with injectable sampling collaborators.""" + + self._detail_sampler = detail_sampler or DetailSampler() + self._tiled_sampling_service = ( + tiled_sampling_service or TiledDiffusionSamplingService() + ) + + def encode(self, vae: Any, pixels: torch.Tensor, tiled: bool) -> Latent: + """Encode pixels into a latent dictionary.""" + + return self._detail_sampler.encode(vae, pixels, tiled) + + def decode(self, vae: Any, latent: Latent, tiled: bool) -> torch.Tensor: + """Decode latent samples into pixels.""" + + return self._detail_sampler.decode(vae, latent, tiled) + + def sample_tiled( + 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 = None, + differential_diffusion: bool = False, + ) -> Latent: + """Sample one latent crop with the selected tiled diffusion runtime.""" + + return self._tiled_sampling_service.sample( + 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, + ) diff --git a/simple_syrup/services/ultralytics_loader_service.py b/simple_syrup/services/ultralytics_loader_service.py new file mode 100644 index 0000000..0158801 --- /dev/null +++ b/simple_syrup/services/ultralytics_loader_service.py @@ -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 + +"""Orchestrate Ultralytics discovery, download, loading, and compatibility.""" + +from __future__ import annotations + +from collections.abc import MutableMapping +from dataclasses import dataclass +from pathlib import Path +from types import ModuleType + +from ..runtime.model_catalog import ULTRALYTICS_ENTRIES, ModelEntry +from ..runtime.model_choices import ModelChoiceService +from ..runtime.model_downloads import DownloadRequest, ModelDownloader, ProgressReporter +from ..runtime.model_folders import expected_model_file, resolve_model_file +from ..runtime.model_instance_cache import ModelInstanceCache +from ..runtime.ultralytics_model_adapter import ( + UltralyticsDetectorModel, + UltralyticsModelAdapter, +) +from ..runtime.ultralytics_model_folders import ( + ULTRALYTICS_BBOX_FOLDER, + ULTRALYTICS_SEGM_FOLDER, + UltralyticsModelFolders, +) +from ..shared.logging import get_logger +from .detector_compat import BBoxDetectorFacade, SegmDetectorFacade + +LOGGER = get_logger(__name__) + +NO_LOCAL_ULTRALYTICS_MODELS = "No local Ultralytics models found" + + +@dataclass(frozen=True) +class LoadedUltralyticsDetector: + """Bundle native and compatibility detector outputs from the loader.""" + + detector_model: UltralyticsDetectorModel + bbox_detector: object + segm_detector: object + + +@dataclass(frozen=True) +class UltralyticsModelCacheKey: + """Identify a loaded Ultralytics detector for process-level reuse.""" + + model_path: Path + + +_LOADED_ULTRALYTICS_MODELS: dict[ + UltralyticsModelCacheKey, LoadedUltralyticsDetector +] = {} + + +class UltralyticsLoaderService: + """Coordinate model selection with runtime adapters and detector facades.""" + + def __init__( + self, + folder_paths_module: ModuleType | None = None, + ultralytics_module: ModuleType | None = None, + downloader: ModelDownloader | None = None, + choice_service: ModelChoiceService | None = None, + cache: ( + MutableMapping[UltralyticsModelCacheKey, LoadedUltralyticsDetector] | None + ) = None, + model_folders: UltralyticsModelFolders | None = None, + model_adapter: UltralyticsModelAdapter | None = None, + ) -> None: + """Create the service with injectable runtime boundaries.""" + + self._model_folders = model_folders or UltralyticsModelFolders( + folder_paths_module + ) + self._model_adapter = model_adapter or UltralyticsModelAdapter( + ultralytics_module + ) + self._downloader = downloader or ModelDownloader() + self._choice_service = choice_service or ModelChoiceService() + self._cache: ModelInstanceCache[ + UltralyticsModelCacheKey, LoadedUltralyticsDetector + ] = ModelInstanceCache( + cache if cache is not None else _LOADED_ULTRALYTICS_MODELS + ) + + def model_choices(self) -> list[str]: + """Return installed choices followed by downloadable catalog choices.""" + + self._model_folders.register() + curated_choices = self._choice_service.ultralytics_choices() + catalog_choice_labels = { + _catalog_selection(entry): entry.display_name + for entry in ULTRALYTICS_ENTRIES + } + available_choices = self.available_models() + visible_catalog_choices = set(curated_choices) + installed_catalog_choices = [ + entry.display_name + for entry in ULTRALYTICS_ENTRIES + if ( + entry.display_name in visible_catalog_choices + and _catalog_selection(entry) in available_choices + ) + ] + installed_non_catalog_choices = [ + choice + for choice in available_choices + if choice not in catalog_choice_labels + ] + downloadable_choices = [ + choice + for choice in curated_choices + if choice not in installed_catalog_choices + ] + choices = ( + installed_non_catalog_choices + + installed_catalog_choices + + downloadable_choices + ) + return choices or [NO_LOCAL_ULTRALYTICS_MODELS] + + def available_models(self) -> list[str]: + """Return model choices discovered by the runtime folder adapter.""" + + return self._model_folders.available_models() + + def load( + self, + model_name: str, + progress: ProgressReporter | None = None, + ) -> LoadedUltralyticsDetector: + """Load one model and construct its workflow-compatible facades.""" + + self.reject_sentinel(model_name) + entry = _catalog_entry_or_none(model_name) + if entry is None: + model_path = self.resolve_model_path(model_name) + normalized_name = model_name.replace("\\", "/") + else: + model_path = self._resolve_catalog_entry(entry, progress) + normalized_name = _catalog_selection(entry) + + key = UltralyticsModelCacheKey(model_path=model_path.resolve()) + already_loaded = key in self._cache.entries + loaded = self._cache.get_or_load( + key, + lambda: self._load_uncached_detector(normalized_name, model_path), + ) + if already_loaded: + LOGGER.info( + "Ultralytics model loaded from process cache", + extra={ + "operation": "load_ultralytics_model", + "model_name": normalized_name, + "model_path": str(model_path), + "task": loaded.detector_model.task, + }, + ) + return loaded + + def reject_sentinel(self, model_name: str) -> None: + """Reject placeholder dropdown selections before filesystem work.""" + + if model_name == NO_LOCAL_ULTRALYTICS_MODELS: + raise ValueError( + "No local Ultralytics models are available. Enable 'Show " + "downloadable models in loader dropdowns' in SimpleSyrup settings " + "or install a model in models\\ultralytics, " + "models\\ultralytics\\bbox, or models\\ultralytics\\segm." + ) + + def resolve_model_path(self, model_name: str) -> Path: + """Resolve a safe model selection through the folder adapter.""" + + return self._model_folders.resolve(model_name) + + def _resolve_catalog_entry( + self, + entry: ModelEntry, + progress: ProgressReporter | None, + ) -> Path: + """Resolve or securely download one curated detector checkpoint.""" + + if len(entry.artifacts) != 1: + raise RuntimeError( + f"Ultralytics catalog entry '{entry.entry_id}' must have one artifact." + ) + self._model_folders.register() + artifact = entry.artifacts[0] + folder_paths_module = self._model_folders.folder_paths_module + existing = resolve_model_file( + artifact.folder_name, + artifact.filename, + folder_paths_module, + ) + destination = existing or expected_model_file( + artifact.folder_name, + artifact.filename, + folder_paths_module, + ) + result = self._downloader.download( + DownloadRequest( + source_url=artifact.source_url, + destination_path=destination, + expected_folder=destination.parent, + description=artifact.description, + expected_sha256=artifact.sha256, + ), + progress, + ) + return result.path + + def _load_uncached_detector( + self, + model_name: str, + model_path: Path, + ) -> LoadedUltralyticsDetector: + """Load a native model and construct both compatibility facades.""" + + detector_model = self._model_adapter.load(model_name, model_path) + bbox_detector = BBoxDetectorFacade(detector_model) + segm_detector = SegmDetectorFacade(detector_model, bbox_detector) + return LoadedUltralyticsDetector( + detector_model=detector_model, + bbox_detector=bbox_detector, + segm_detector=segm_detector, + ) + + +def _catalog_entry_or_none(selection: str) -> ModelEntry | None: + """Return a curated Ultralytics entry when a dropdown label matches it.""" + + return next( + ( + entry + for entry in ULTRALYTICS_ENTRIES + if selection in (entry.entry_id, entry.display_name) + ), + None, + ) + + +def _catalog_selection(entry: ModelEntry) -> str: + """Return the local conventional selection path for one catalog entry.""" + + if len(entry.artifacts) != 1: + raise ValueError( + f"Ultralytics catalog entry '{entry.entry_id}' must have one artifact." + ) + artifact = entry.artifacts[0] + prefix_by_folder = { + ULTRALYTICS_BBOX_FOLDER: "bbox", + ULTRALYTICS_SEGM_FOLDER: "segm", + } + try: + prefix = prefix_by_folder[artifact.folder_name] + except KeyError as error: + raise ValueError( + f"Ultralytics catalog entry '{entry.entry_id}' has unsupported folder " + f"'{artifact.folder_name}'." + ) from error + return f"{prefix}/{artifact.filename}" diff --git a/tests/AGENTS.md b/tests/AGENTS.md new file mode 100644 index 0000000..3b4c1b9 --- /dev/null +++ b/tests/AGENTS.md @@ -0,0 +1,100 @@ +# AGENTS.md + +## Scope + +This file supplements the repository-root `AGENTS.md` for every file below +`tests/`. + +## Mission + +The test suite provides fast, deterministic, behaviorally meaningful evidence +for SimpleSyrup's ComfyUI nodes, domain behavior, services, runtime adapters, +tooling, and persisted workflow contracts. + +## Ownership And Placement + +- Organize tests by product capability or authoritative behavior owner. +- Do not add test modules directly under `tests/`; the root is reserved for + pytest configuration and execution policy. +- Within a large capability, split tests by domain, service, runtime, node, or + integration boundary when those owners change independently. +- Keep capability-specific fixtures, fakes, values, and harnesses in that + capability's `support/` package. +- Promote support to `tests/support/` only when independent capabilities use + the same stable testing contract. +- Do not create generic helper, common, misc, or utility dumping grounds. +- Keep `web/tests/` organized by the same capabilities as `web/src/`. +- Update policy entries, runner inventories, imports, and focused-test paths in + the same change as any test move. + +## Behavioral Proof + +- Test observable behavior through its authoritative owner. +- Prefer the lightest real component that proves the complete contract. +- Do not mock the behavior under test or duplicate production rules in expected + result implementations. +- Cover relevant success, failure, boundary, cancellation, cleanup, and + regression paths. +- System and live-Comfy tests must prove composition that focused owner tests + cannot prove. + +## Parallelism And Isolation + +- Tests are parallel-safe by default. +- Do not branch behavior on `PYTEST_XDIST_WORKER`. +- Do not serialize tests to hide leaked state, fixed resources, nondeterminism, + or unsafe cleanup. +- Use fresh-process isolation only for a demonstrated process-lifetime + constraint. +- Use serial execution only for an exact global or external resource that + cannot safely overlap. +- Record every isolated or serial module in `tests/ci_test_policy.py` with a + reviewed governance disposition. + +## Determinism + +- Control clocks, timers, randomness, environment, filesystems, subprocesses, + network responses, and external Comfy state when they affect behavior. +- Do not use arbitrary sleeps as completion conditions. +- Wait for observable state with bounded diagnostic timeouts. +- Do not use retries, skips, weakened assertions, or increased delays to hide + flakes. +- Every subprocess and network operation must have an explicit failure bound + and guaranteed cleanup. + +## Test State + +- Use `tmp_path` and `pathlib.Path` for filesystem behavior. +- Do not write test artifacts into the repository unless the artifact path is + itself the contract under test. +- Restore environment variables, module replacements, registries, logging + handlers, working directories, and global settings. +- Keep autouse fixtures limited to universal safety and cleanup invariants. +- Give every fixture one cohesive lifecycle and explicit typed result. + +## Typing And Maintainability + +- Type tests, fixtures, fakes, builders, and harness APIs. +- Do not add test-only `.pyi` files that shadow executable modules. +- Use explicit protocols or focused typed fakes at dynamic boundaries. +- Keep setup and assertions readable at the test callsite. +- Shared abstractions must reduce repeated change risk without hiding the + behavior being proved. + +## Governance + +- Run `..\..\venv\Scripts\python.exe -m tools.check_test_governance` + after changing test placement, timing, isolation, resources, or execution + policy. +- Every discovered candidate requires source-level review. +- A classification waiver records legitimate intentional behavior. +- Inappropriate current design requires debt plus an exact remediation waiver. +- Reviewed state must remain fingerprinted, expiring, and file-specific. + +## Verification + +- Run focused tests continuously while changing a capability. +- Run collection after moving tests. +- Run the full parallel suite before completion. +- Run architecture governance, test governance, formatting, lint, strict + typing, and frontend gates when applicable. diff --git a/tests/ci_test_policy.py b/tests/ci_test_policy.py new file mode 100644 index 0000000..eee26c5 --- /dev/null +++ b/tests/ci_test_policy.py @@ -0,0 +1,10 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Declare reviewed process-isolated and serial test modules.""" + +from __future__ import annotations + +ISOLATED_TEST_MODULES: tuple[str, ...] = () +SERIAL_TEST_MODULES: tuple[str, ...] = () diff --git a/tests/comfy_integration/__init__.py b/tests/comfy_integration/__init__.py new file mode 100644 index 0000000..dd902da --- /dev/null +++ b/tests/comfy_integration/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own comfy integration test behavior.""" diff --git a/tests/test_comfy_adapter_evidence.py b/tests/comfy_integration/test_comfy_adapter_evidence.py similarity index 100% rename from tests/test_comfy_adapter_evidence.py rename to tests/comfy_integration/test_comfy_adapter_evidence.py diff --git a/tests/test_comfy_adapter_identity_index.py b/tests/comfy_integration/test_comfy_adapter_identity_index.py similarity index 100% rename from tests/test_comfy_adapter_identity_index.py rename to tests/comfy_integration/test_comfy_adapter_identity_index.py diff --git a/tests/test_comfy_api.py b/tests/comfy_integration/test_comfy_api.py similarity index 100% rename from tests/test_comfy_api.py rename to tests/comfy_integration/test_comfy_api.py diff --git a/tests/test_comfy_conditioning_model_loader.py b/tests/comfy_integration/test_comfy_conditioning_model_loader.py similarity index 100% rename from tests/test_comfy_conditioning_model_loader.py rename to tests/comfy_integration/test_comfy_conditioning_model_loader.py diff --git a/tests/test_comfy_execution_trace.py b/tests/comfy_integration/test_comfy_execution_trace.py similarity index 100% rename from tests/test_comfy_execution_trace.py rename to tests/comfy_integration/test_comfy_execution_trace.py diff --git a/tests/test_comfy_integration_artifacts.py b/tests/comfy_integration/test_comfy_integration_artifacts.py similarity index 100% rename from tests/test_comfy_integration_artifacts.py rename to tests/comfy_integration/test_comfy_integration_artifacts.py diff --git a/tests/test_comfy_integration_baseline_workflow.py b/tests/comfy_integration/test_comfy_integration_baseline_workflow.py similarity index 100% rename from tests/test_comfy_integration_baseline_workflow.py rename to tests/comfy_integration/test_comfy_integration_baseline_workflow.py diff --git a/tests/test_comfy_integration_history_output.py b/tests/comfy_integration/test_comfy_integration_history_output.py similarity index 100% rename from tests/test_comfy_integration_history_output.py rename to tests/comfy_integration/test_comfy_integration_history_output.py diff --git a/tests/test_comfy_integration_loopback_port.py b/tests/comfy_integration/test_comfy_integration_loopback_port.py similarity index 69% rename from tests/test_comfy_integration_loopback_port.py rename to tests/comfy_integration/test_comfy_integration_loopback_port.py index 6f882b3..efdc15f 100644 --- a/tests/test_comfy_integration_loopback_port.py +++ b/tests/comfy_integration/test_comfy_integration_loopback_port.py @@ -12,17 +12,20 @@ import pytest from tools.comfy_integration.loopback_port import ( is_loopback_port_available, - select_unused_loopback_port, + reserve_loopback_port, validate_loopback_port, ) -def test_selected_port_is_nondefault_and_unused() -> None: - """Prove selection returns a bindable non-default loopback port.""" +def test_reserved_port_remains_exclusive_until_owner_releases_it() -> None: + """Keep the OS-assigned port unavailable throughout reservation ownership.""" - port = select_unused_loopback_port() + with reserve_loopback_port() as reservation: + port = reservation.port + assert port not in {8188, 8297} + assert not is_loopback_port_available(port) - assert port not in {8188, 8297} + assert is_loopback_port_available(port) with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: probe.bind(("127.0.0.1", port)) assert not is_loopback_port_available(port) diff --git a/tests/test_comfy_integration_managed_server.py b/tests/comfy_integration/test_comfy_integration_managed_server.py similarity index 66% rename from tests/test_comfy_integration_managed_server.py rename to tests/comfy_integration/test_comfy_integration_managed_server.py index 5041703..40e0608 100644 --- a/tests/test_comfy_integration_managed_server.py +++ b/tests/comfy_integration/test_comfy_integration_managed_server.py @@ -52,6 +52,21 @@ class _FakeClient: self.base_url = base_url +class _FakeReservation: + """Expose one deterministic candidate port and release observation.""" + + def __init__(self, port: int) -> None: + """Retain the selected fake port.""" + + self.port = port + self.release_calls = 0 + + def release(self) -> None: + """Record the handoff immediately before process launch.""" + + self.release_calls += 1 + + def _configure( monkeypatch: pytest.MonkeyPatch, process: _FakeProcess, @@ -60,7 +75,8 @@ def _configure( ) -> None: """Replace external lifecycle boundaries with deterministic fakes.""" - monkeypatch.setattr(managed_server, "select_unused_loopback_port", lambda: 8299) + reservation = _FakeReservation(8299) + monkeypatch.setattr(managed_server, "reserve_loopback_port", lambda: reservation) monkeypatch.setattr( "tools.comfy_integration.managed_server.WindowsComfyProcess.start", lambda command, stdout_path, stderr_path: process, @@ -80,6 +96,61 @@ def _configure( monkeypatch.setattr(managed_server, "wait_for_server", ready) +def test_port_collision_retries_with_new_owned_process( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Retry only after cleanup proves another process claimed the candidate port.""" + + reservations = [_FakeReservation(8299), _FakeReservation(8300)] + first_reservation, second_reservation = reservations + processes = [_FakeProcess(), _FakeProcess()] + monkeypatch.setattr( + managed_server, + "reserve_loopback_port", + lambda: reservations.pop(0), + ) + monkeypatch.setattr( + "tools.comfy_integration.managed_server.WindowsComfyProcess.start", + lambda command, stdout_path, stderr_path: processes.pop(0), + ) + monkeypatch.setattr(managed_server, "LoopbackComfyClient", _FakeClient) + readiness_calls = 0 + + def ready( + client: object, live: object, required: object, *, timeout: float + ) -> JsonObject: + """Fail the collided attempt and accept the replacement.""" + + nonlocal readiness_calls + del client, live, required, timeout + readiness_calls += 1 + if readiness_calls == 1: + raise RuntimeError("Managed Comfy process exited before readiness.") + return {"ready": True} + + monkeypatch.setattr(managed_server, "wait_for_server", ready) + availability = iter((False,)) + monkeypatch.setattr( + managed_server, + "is_loopback_port_available", + lambda _port: next(availability), + ) + first, second = processes + + with ManagedComfyServer( + comfy_root=Path(""), + artifacts=IntegrationArtifacts(tmp_path), + required_node_ids=frozenset(), + ) as running: + assert running.port == 8300 + + assert first.stop_calls == 1 + assert second.stop_calls == 1 + assert first_reservation.release_calls == 1 + assert second_reservation.release_calls == 1 + assert readiness_calls == 2 + + def test_context_stops_exact_created_process_after_body_failure( tmp_path: Path, monkeypatch: pytest.MonkeyPatch ) -> None: diff --git a/tests/test_comfy_integration_readiness.py b/tests/comfy_integration/test_comfy_integration_readiness.py similarity index 100% rename from tests/test_comfy_integration_readiness.py rename to tests/comfy_integration/test_comfy_integration_readiness.py diff --git a/tests/test_comfy_integration_server_process.py b/tests/comfy_integration/test_comfy_integration_server_process.py similarity index 100% rename from tests/test_comfy_integration_server_process.py rename to tests/comfy_integration/test_comfy_integration_server_process.py diff --git a/tests/test_comfy_latent_normalization.py b/tests/comfy_integration/test_comfy_latent_normalization.py similarity index 100% rename from tests/test_comfy_latent_normalization.py rename to tests/comfy_integration/test_comfy_latent_normalization.py diff --git a/tests/test_comfy_regional_adapter_resolver.py b/tests/comfy_integration/test_comfy_regional_adapter_resolver.py similarity index 100% rename from tests/test_comfy_regional_adapter_resolver.py rename to tests/comfy_integration/test_comfy_regional_adapter_resolver.py diff --git a/tests/test_comfy_regional_conditioning_processing.py b/tests/comfy_integration/test_comfy_regional_conditioning_processing.py similarity index 99% rename from tests/test_comfy_regional_conditioning_processing.py rename to tests/comfy_integration/test_comfy_regional_conditioning_processing.py index b2fba77..8b2fc31 100644 --- a/tests/test_comfy_regional_conditioning_processing.py +++ b/tests/comfy_integration/test_comfy_regional_conditioning_processing.py @@ -6,7 +6,6 @@ from __future__ import annotations -from pathlib import Path from types import SimpleNamespace from typing import Any, cast from uuid import UUID @@ -14,7 +13,11 @@ from uuid import UUID import comfy.conds import pytest import torch +from support.repository import REPOSITORY_ROOT +from simple_syrup.domain.attention_coupling_preparation import ( + AttentionCouplingPreparation, +) from simple_syrup.domain.conditioning_batch import ConditioningBatch from simple_syrup.domain.raw_regional_attention import ( build_raw_regional_attention_plan, @@ -37,7 +40,6 @@ from simple_syrup.runtime.ppm_negpip_interop import ( ) from simple_syrup.services.attention_coupling_preparation_service import ( ATTENTION_COUPLING_PREPARATION_SERVICE, - AttentionCouplingPreparation, ) @@ -370,7 +372,7 @@ def test_processing_owner_contains_no_semantic_token_repeat_path() -> None: """Keep semantic-length rejection explicit in the production source.""" source_path = ( - Path(__file__).resolve().parents[1] + REPOSITORY_ROOT / "simple_syrup" / "runtime" / "attention_coupling" diff --git a/tests/test_comfy_tool_defaults.py b/tests/comfy_integration/test_comfy_tool_defaults.py similarity index 100% rename from tests/test_comfy_tool_defaults.py rename to tests/comfy_integration/test_comfy_tool_defaults.py diff --git a/tests/test_diffusion_wrapper_invocation.py b/tests/comfy_integration/test_diffusion_wrapper_invocation.py similarity index 100% rename from tests/test_diffusion_wrapper_invocation.py rename to tests/comfy_integration/test_diffusion_wrapper_invocation.py diff --git a/tests/test_effective_model_graph.py b/tests/comfy_integration/test_effective_model_graph.py similarity index 100% rename from tests/test_effective_model_graph.py rename to tests/comfy_integration/test_effective_model_graph.py diff --git a/tests/test_graph_provenance.py b/tests/comfy_integration/test_graph_provenance.py similarity index 100% rename from tests/test_graph_provenance.py rename to tests/comfy_integration/test_graph_provenance.py diff --git a/tests/conftest.py b/tests/conftest.py index 4143988..1242259 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -17,10 +17,11 @@ if not CUDA_TESTS_ENABLED: os.environ["CUDA_VISIBLE_DEVICES"] = "-1" PROJECT_ROOT = Path(__file__).resolve().parents[1] +TESTS_ROOT = PROJECT_ROOT / "tests" COMFY_ROOT = PROJECT_ROOT.parents[1] CUSTOM_NODES_ROOT = PROJECT_ROOT.parent -for path in (PROJECT_ROOT, COMFY_ROOT, CUSTOM_NODES_ROOT): +for path in reversed((TESTS_ROOT, PROJECT_ROOT, COMFY_ROOT, CUSTOM_NODES_ROOT)): path_text = str(path) if path_text not in sys.path: sys.path.insert(0, path_text) diff --git a/tests/crosscutting/__init__.py b/tests/crosscutting/__init__.py new file mode 100644 index 0000000..9904676 --- /dev/null +++ b/tests/crosscutting/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own crosscutting test behavior.""" diff --git a/tests/crosscutting/contracts/__init__.py b/tests/crosscutting/contracts/__init__.py new file mode 100644 index 0000000..6306a36 --- /dev/null +++ b/tests/crosscutting/contracts/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own crosscutting contracts test behavior.""" diff --git a/tests/test_license_headers.py b/tests/crosscutting/contracts/test_license_headers.py similarity index 97% rename from tests/test_license_headers.py rename to tests/crosscutting/contracts/test_license_headers.py index 9c418a1..3f14f9f 100644 --- a/tests/test_license_headers.py +++ b/tests/crosscutting/contracts/test_license_headers.py @@ -12,7 +12,9 @@ from pathlib import Path from types import ModuleType from typing import Any, cast -REPO_ROOT = Path(__file__).resolve().parents[1] +from support.repository import REPOSITORY_ROOT + +REPO_ROOT = REPOSITORY_ROOT TOOLS_MODULE = REPO_ROOT / "tools" / "add_license_headers.py" PROJECT_LINE = "SimpleSyrup - workflow-focused ComfyUI extensions for image generation" SPDX_LINE = "SPDX-License-Identifier: AGPL-3.0-or-later" diff --git a/tests/test_no_external_pack_imports.py b/tests/crosscutting/contracts/test_no_external_pack_imports.py similarity index 91% rename from tests/test_no_external_pack_imports.py rename to tests/crosscutting/contracts/test_no_external_pack_imports.py index 97d14e0..78cdcdc 100644 --- a/tests/test_no_external_pack_imports.py +++ b/tests/crosscutting/contracts/test_no_external_pack_imports.py @@ -6,13 +6,13 @@ from __future__ import annotations -from pathlib import Path +from support.repository import REPOSITORY_ROOT def test_simple_syrup_does_not_import_impact_or_layerstyle() -> None: """Runtime source should not import Impact Pack or LayerStyle modules.""" - project_root = Path(__file__).resolve().parents[1] + project_root = REPOSITORY_ROOT source_text = "\n".join( path.read_text(encoding="utf-8") for path in (project_root / "simple_syrup").rglob("*.py") diff --git a/tests/test_node_tooltips.py b/tests/crosscutting/contracts/test_node_tooltips.py similarity index 100% rename from tests/test_node_tooltips.py rename to tests/crosscutting/contracts/test_node_tooltips.py diff --git a/tests/test_packaging_metadata.py b/tests/crosscutting/contracts/test_packaging_metadata.py similarity index 98% rename from tests/test_packaging_metadata.py rename to tests/crosscutting/contracts/test_packaging_metadata.py index 060d878..d5f26fd 100644 --- a/tests/test_packaging_metadata.py +++ b/tests/crosscutting/contracts/test_packaging_metadata.py @@ -12,8 +12,9 @@ from importlib import import_module from pathlib import Path from packaging.requirements import Requirement +from support.repository import REPOSITORY_ROOT -REPO_ROOT = Path(__file__).resolve().parents[1] +REPO_ROOT = REPOSITORY_ROOT COMFY_ROOT = REPO_ROOT.parents[1] EXPECTED_RUNTIME_REQUIREMENTS = ( "torchlanc", diff --git a/tests/test_persisted_widget_order_contract.py b/tests/crosscutting/contracts/test_persisted_widget_order_contract.py similarity index 100% rename from tests/test_persisted_widget_order_contract.py rename to tests/crosscutting/contracts/test_persisted_widget_order_contract.py diff --git a/tests/test_registration.py b/tests/crosscutting/contracts/test_registration.py similarity index 94% rename from tests/test_registration.py rename to tests/crosscutting/contracts/test_registration.py index bf2e6b9..91533a0 100644 --- a/tests/test_registration.py +++ b/tests/crosscutting/contracts/test_registration.py @@ -10,11 +10,11 @@ import asyncio import importlib import subprocess import sys -from pathlib import Path from types import ModuleType from typing import Any, Protocol, cast import pytest +from support.repository import REPOSITORY_ROOT BASE_NODE_IDS = [ "SimpleSyrup.AllPromptAttentionSEGS", @@ -117,7 +117,7 @@ def test_nodes_package_is_not_a_legacy_registry() -> None: def test_package_imports_from_custom_nodes_parent_path() -> None: """ComfyUI-style import works without the repository root on sys.path.""" - project_root = Path(__file__).resolve().parents[1] + project_root = REPOSITORY_ROOT custom_nodes_root = project_root.parent script = ( "import importlib, pathlib, sys; " @@ -138,6 +138,7 @@ def test_package_imports_from_custom_nodes_parent_path() -> None: check=False, capture_output=True, text=True, + timeout=60.0, ) assert result.returncode == 0, result.stderr @@ -146,7 +147,7 @@ def test_package_imports_from_custom_nodes_parent_path() -> None: def test_comfy_import_exposes_stable_internal_package_alias() -> None: """ComfyUI-style import exposes `simple_syrup` to nested vendored packages.""" - project_root = Path(__file__).resolve().parents[1] + project_root = REPOSITORY_ROOT custom_nodes_root = project_root.parent script = ( "import importlib, pathlib, sys; " @@ -167,15 +168,18 @@ def test_comfy_import_exposes_stable_internal_package_alias() -> None: check=False, capture_output=True, text=True, + timeout=60.0, ) assert result.returncode == 0, result.stderr -def test_registration_import_does_not_require_torchlanc() -> None: +def test_registration_import_does_not_require_torchlanc( + monkeypatch: pytest.MonkeyPatch, +) -> None: """Importing registration does not eagerly import TorchLanc.""" - sys.modules.pop("torchlanc", None) + monkeypatch.delitem(sys.modules, "torchlanc", raising=False) importlib.import_module("SimpleSyrup") imported_module: ModuleType | None = sys.modules.get("torchlanc") @@ -187,7 +191,7 @@ def test_v3_entrypoint_exports_all_base_nodes_without_prompt_control( ) -> None: """Comfy v3 entrypoint exports every maintained non-conditional node.""" - sys.modules.pop("prompt_control.nodes_lazy", None) + monkeypatch.delitem(sys.modules, "prompt_control.nodes_lazy", raising=False) package = importlib.import_module("SimpleSyrup") nodes_v3 = importlib.import_module("SimpleSyrup.simple_syrup.nodes_v3") monkeypatch.setattr(nodes_v3, "prompt_control_is_available", lambda: False) @@ -204,7 +208,7 @@ def test_v3_entrypoint_adds_only_prompt_control_nodes_when_available( ) -> None: """Prompt Control availability adds conditional nodes without removing others.""" - sys.modules.pop("prompt_control.nodes_lazy", None) + monkeypatch.delitem(sys.modules, "prompt_control.nodes_lazy", raising=False) 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) diff --git a/tests/test_repository_path_portability.py b/tests/crosscutting/contracts/test_repository_path_portability.py similarity index 93% rename from tests/test_repository_path_portability.py rename to tests/crosscutting/contracts/test_repository_path_portability.py index 2cf6017..06023fb 100644 --- a/tests/test_repository_path_portability.py +++ b/tests/crosscutting/contracts/test_repository_path_portability.py @@ -7,9 +7,10 @@ from __future__ import annotations import re -from pathlib import Path -REPOSITORY_ROOT = Path(__file__).resolve().parents[1] +from support.repository import REPOSITORY_ROOT + +REPOSITORY_ROOT = REPOSITORY_ROOT EXECUTABLE_ROOTS = ( REPOSITORY_ROOT / "simple_syrup", REPOSITORY_ROOT / "tools", diff --git a/tests/test_third_party_vendoring_contract.py b/tests/crosscutting/contracts/test_third_party_vendoring_contract.py similarity index 99% rename from tests/test_third_party_vendoring_contract.py rename to tests/crosscutting/contracts/test_third_party_vendoring_contract.py index 01b6653..111bfb2 100644 --- a/tests/test_third_party_vendoring_contract.py +++ b/tests/crosscutting/contracts/test_third_party_vendoring_contract.py @@ -9,7 +9,9 @@ from __future__ import annotations import tomllib from pathlib import Path -REPO_ROOT = Path(__file__).resolve().parents[1] +from support.repository import REPOSITORY_ROOT + +REPO_ROOT = REPOSITORY_ROOT def test_third_party_manifest_references_existing_licenses_and_runtime_files() -> None: diff --git a/tests/crosscutting/governance/__init__.py b/tests/crosscutting/governance/__init__.py new file mode 100644 index 0000000..8f821a3 --- /dev/null +++ b/tests/crosscutting/governance/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own executable repository-governance proof.""" diff --git a/tests/crosscutting/governance/test_architecture_governance.py b/tests/crosscutting/governance/test_architecture_governance.py new file mode 100644 index 0000000..30406d7 --- /dev/null +++ b/tests/crosscutting/governance/test_architecture_governance.py @@ -0,0 +1,99 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Verify structural governance and the current architecture inventory.""" + +from __future__ import annotations + +from datetime import date +from pathlib import Path + +from support.repository import REPOSITORY_ROOT + +from tools.architecture_governance.import_boundaries import validate_import_boundaries +from tools.architecture_governance.metrics import production_line_count +from tools.architecture_governance.soft_reviews import validate_soft_reviews +from tools.architecture_governance.validation import validate_repository + + +def _write(path: Path, content: str) -> None: + """Write one isolated governance fixture.""" + + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(content, encoding="utf-8") + + +def _write_policy(root: Path) -> None: + """Create a minimal strict architecture policy.""" + + _write( + root / "governance/architecture/policy.toml", + """schema_version = 2 +[structure] +soft_lines = 2 +hard_lines = 4 +source_roots = ["product"] +source_files = [] +source_extensions = [".py", ".ts"] +excluded_paths = [] +[registries] +debt = "governance/architecture/debt.toml" +waivers = "governance/architecture/waivers.toml" +""", + ) + _write( + root / "governance/architecture/debt.toml", + "schema_version = 1\ndebts = []\n", + ) + _write( + root / "governance/architecture/waivers.toml", + "schema_version = 1\nwaivers = []\n", + ) + + +def test_unclassified_hard_overage_is_blocking(tmp_path: Path) -> None: + """Require human ownership assessment for every hard-gate file.""" + + _write_policy(tmp_path) + _write( + tmp_path / "product/mixed.py", + "\n".join(f"VALUE_{index} = {index}" for index in range(6)), + ) + + diagnostics = validate_repository(tmp_path, today=date(2026, 9, 24)) + + assert any(item.rule == "STRUCT003" for item in diagnostics) + + +def test_typescript_comments_do_not_inflate_production_lines(tmp_path: Path) -> None: + """Count authored TypeScript while excluding standalone comments.""" + + path = tmp_path / "source.ts" + _write(path, "// heading\n/* block\ncomment */\nconst value = 1;\n") + + assert production_line_count(path) == 1 + + +def test_current_repository_has_no_architecture_governance_errors() -> None: + """Keep every hard structural finding exactly reviewed.""" + + errors = [ + item + for item in validate_repository(REPOSITORY_ROOT) + if item.severity == "error" + ] + + assert errors == [] + + +def test_current_repository_has_no_unreviewed_import_boundary_debt() -> None: + """Keep every dependency-direction violation exactly inventoried.""" + + assert validate_import_boundaries(REPOSITORY_ROOT) == [] + + +def test_current_repository_has_exact_soft_ceiling_reviews() -> None: + """Keep every advisory structural warning source-reviewed and current.""" + + assert validate_soft_reviews(REPOSITORY_ROOT) == [] diff --git a/tests/crosscutting/governance/test_test_governance.py b/tests/crosscutting/governance/test_test_governance.py new file mode 100644 index 0000000..7684592 --- /dev/null +++ b/tests/crosscutting/governance/test_test_governance.py @@ -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 + +"""Verify capability placement and reviewed test-governance state.""" + +from __future__ import annotations + +from pathlib import Path + +from support.repository import REPOSITORY_ROOT + +from tools.test_governance.discovery import ( + LAYOUT_RULE, + TYPESCRIPT_LAYOUT_RULE, + TYPESCRIPT_OPTIONAL_RULE, + discover_test_candidates, +) +from tools.test_governance.loading import load_test_policy +from tools.test_governance.validation import validate_test_governance + + +def _write(path: Path, content: str) -> None: + """Write one isolated test-governance fixture.""" + + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(content, encoding="utf-8") + + +def _write_fixture(root: Path) -> None: + """Create the smallest complete test-governance repository.""" + + _write( + root / "governance/testing/policy.toml", + """schema_version = 1 +[scope] +test_root = "tests" +semantic_support_roots = ["tools"] +root_source_extensions = [".py", ".pyi"] +allowed_root_source_paths = ["tests/conftest.py", "tests/ci_test_policy.py"] +[discovery] +serial_policy = "tests/ci_test_policy.py" +wait_calls = ["QTest.qWait", "time.sleep"] +wall_clock_calls = ["time.perf_counter"] +xdist_environment_name = "PYTEST_XDIST_WORKER" +repository_scratch_name = ".pytest-tmp" +[registries] +debt = "governance/testing/debt.toml" +waivers = "governance/testing/waivers.toml" +""", + ) + _write(root / "governance/testing/debt.toml", "schema_version = 1\ndebts = []\n") + _write( + root / "governance/testing/waivers.toml", + "schema_version = 1\nwaivers = []\n", + ) + _write( + root / "tests/ci_test_policy.py", + "ISOLATED_TEST_MODULES = ()\nSERIAL_TEST_MODULES = ()\n", + ) + _write(root / "tests/conftest.py", "") + _write(root / "tools/__init__.py", "") + + +def test_discovery_covers_python_and_frontend_root_layout(tmp_path: Path) -> None: + """Reject authored tests left outside capability owners.""" + + _write_fixture(tmp_path) + _write(tmp_path / "tests/test_root.py", "def test_root() -> None: pass\n") + _write(tmp_path / "web/tests/root.test.ts", "test('root', () => {});\n") + policy = load_test_policy(tmp_path / "governance/testing/policy.toml") + + rules = {item.rule for item in discover_test_candidates(tmp_path, policy)} + + assert {LAYOUT_RULE, TYPESCRIPT_LAYOUT_RULE} <= rules + + +def test_discovery_reports_optional_frontend_proof(tmp_path: Path) -> None: + """Require review of skipped, todo, or exclusive frontend proof.""" + + _write_fixture(tmp_path) + _write( + tmp_path / "web/tests/media/preview.test.ts", + "describe.only('preview', () => {});\n", + ) + policy = load_test_policy(tmp_path / "governance/testing/policy.toml") + + rules = {item.rule for item in discover_test_candidates(tmp_path, policy)} + + assert TYPESCRIPT_OPTIONAL_RULE in rules + + +def test_current_repository_has_exact_reviewed_test_state() -> None: + """Keep every reliability candidate exactly reviewed.""" + + result = validate_test_governance(REPOSITORY_ROOT) + + assert not [item for item in result.diagnostics if item.severity == "error"] diff --git a/tests/external_llm/__init__.py b/tests/external_llm/__init__.py new file mode 100644 index 0000000..bad5d54 --- /dev/null +++ b/tests/external_llm/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own external llm test behavior.""" diff --git a/tests/test_external_llm_client.py b/tests/external_llm/test_external_llm_client.py similarity index 100% rename from tests/test_external_llm_client.py rename to tests/external_llm/test_external_llm_client.py diff --git a/tests/test_external_llm_images.py b/tests/external_llm/test_external_llm_images.py similarity index 100% rename from tests/test_external_llm_images.py rename to tests/external_llm/test_external_llm_images.py diff --git a/tests/test_external_llm_keyring.py b/tests/external_llm/test_external_llm_keyring.py similarity index 100% rename from tests/test_external_llm_keyring.py rename to tests/external_llm/test_external_llm_keyring.py diff --git a/tests/test_external_llm_prompt_node.py b/tests/external_llm/test_external_llm_prompt_node.py similarity index 100% rename from tests/test_external_llm_prompt_node.py rename to tests/external_llm/test_external_llm_prompt_node.py diff --git a/tests/test_external_llm_prompt_service.py b/tests/external_llm/test_external_llm_prompt_service.py similarity index 100% rename from tests/test_external_llm_prompt_service.py rename to tests/external_llm/test_external_llm_prompt_service.py diff --git a/tests/test_external_llm_prompt_v3_node.py b/tests/external_llm/test_external_llm_prompt_v3_node.py similarity index 100% rename from tests/test_external_llm_prompt_v3_node.py rename to tests/external_llm/test_external_llm_prompt_v3_node.py diff --git a/tests/test_external_llm_routes.py b/tests/external_llm/test_external_llm_routes.py similarity index 99% rename from tests/test_external_llm_routes.py rename to tests/external_llm/test_external_llm_routes.py index e81e697..4c301e0 100644 --- a/tests/test_external_llm_routes.py +++ b/tests/external_llm/test_external_llm_routes.py @@ -18,7 +18,7 @@ from simple_syrup.domain.external_llm import ( ExternalLLMChatResponse, ExternalLLMProviderError, ) -from simple_syrup.runtime.external_llm_routes import ( +from simple_syrup.integration.external_llm_routes import ( EXTERNAL_LLM_API_KEY_ROUTE, EXTERNAL_LLM_MODELS_REFRESH_ROUTE, EXTERNAL_LLM_SETTINGS_ROUTE, diff --git a/tests/test_external_llm_segs_images.py b/tests/external_llm/test_external_llm_segs_images.py similarity index 100% rename from tests/test_external_llm_segs_images.py rename to tests/external_llm/test_external_llm_segs_images.py diff --git a/tests/media/__init__.py b/tests/media/__init__.py new file mode 100644 index 0000000..4fa4df8 --- /dev/null +++ b/tests/media/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own media test behavior.""" diff --git a/tests/test_image_file_loader.py b/tests/media/test_image_file_loader.py similarity index 100% rename from tests/test_image_file_loader.py rename to tests/media/test_image_file_loader.py diff --git a/tests/test_load_image_list_service.py b/tests/media/test_load_image_list_service.py similarity index 100% rename from tests/test_load_image_list_service.py rename to tests/media/test_load_image_list_service.py diff --git a/tests/test_load_image_list_v3_node.py b/tests/media/test_load_image_list_v3_node.py similarity index 100% rename from tests/test_load_image_list_v3_node.py rename to tests/media/test_load_image_list_v3_node.py diff --git a/tests/test_load_mask_batch_service.py b/tests/media/test_load_mask_batch_service.py similarity index 100% rename from tests/test_load_mask_batch_service.py rename to tests/media/test_load_mask_batch_service.py diff --git a/tests/test_load_mask_batch_v3_node.py b/tests/media/test_load_mask_batch_v3_node.py similarity index 100% rename from tests/test_load_mask_batch_v3_node.py rename to tests/media/test_load_mask_batch_v3_node.py diff --git a/tests/test_mask_batch_preview_routes.py b/tests/media/test_mask_batch_preview_routes.py similarity index 98% rename from tests/test_mask_batch_preview_routes.py rename to tests/media/test_mask_batch_preview_routes.py index 99f8e3a..6c40304 100644 --- a/tests/test_mask_batch_preview_routes.py +++ b/tests/media/test_mask_batch_preview_routes.py @@ -17,8 +17,8 @@ import pytest from aiohttp import web from PIL import Image -import simple_syrup.runtime.mask_batch_preview_routes as preview_routes -from simple_syrup.runtime.mask_batch_preview_routes import ( +import simple_syrup.integration.mask_batch_preview_routes as preview_routes +from simple_syrup.integration.mask_batch_preview_routes import ( MASK_BATCH_PREVIEW_ROUTE, Handler, MaskBatchLoaderProtocol, diff --git a/tests/test_mask_file_loader.py b/tests/media/test_mask_file_loader.py similarity index 100% rename from tests/test_mask_file_loader.py rename to tests/media/test_mask_file_loader.py diff --git a/tests/test_ordered_files.py b/tests/media/test_ordered_files.py similarity index 100% rename from tests/test_ordered_files.py rename to tests/media/test_ordered_files.py diff --git a/tests/test_resize_geometry.py b/tests/media/test_resize_geometry.py similarity index 100% rename from tests/test_resize_geometry.py rename to tests/media/test_resize_geometry.py diff --git a/tests/test_resize_node.py b/tests/media/test_resize_node.py similarity index 100% rename from tests/test_resize_node.py rename to tests/media/test_resize_node.py diff --git a/tests/test_resize_resamplers.py b/tests/media/test_resize_resamplers.py similarity index 100% rename from tests/test_resize_resamplers.py rename to tests/media/test_resize_resamplers.py diff --git a/tests/test_resize_service.py b/tests/media/test_resize_service.py similarity index 100% rename from tests/test_resize_service.py rename to tests/media/test_resize_service.py diff --git a/tests/test_scale_factor_node.py b/tests/media/test_scale_factor_node.py similarity index 100% rename from tests/test_scale_factor_node.py rename to tests/media/test_scale_factor_node.py diff --git a/tests/test_scale_factor_v3_node.py b/tests/media/test_scale_factor_v3_node.py similarity index 100% rename from tests/test_scale_factor_v3_node.py rename to tests/media/test_scale_factor_v3_node.py diff --git a/tests/test_seg_preview_assets.py b/tests/media/test_seg_preview_assets.py similarity index 98% rename from tests/test_seg_preview_assets.py rename to tests/media/test_seg_preview_assets.py index bb62d3f..da2b620 100644 --- a/tests/test_seg_preview_assets.py +++ b/tests/media/test_seg_preview_assets.py @@ -13,13 +13,13 @@ import pytest import torch from pytest import MonkeyPatch -from simple_syrup.domain.segs import CropRegion -from simple_syrup.runtime.seg_preview_assets import ComfySegPreviewAssetPublisher -from simple_syrup.services.simple_preview_segs_service import ( +from simple_syrup.domain.seg_preview import ( AtlasPlacement, SegPreviewDocument, SegPreviewRegion, ) +from simple_syrup.domain.segs import CropRegion +from simple_syrup.runtime.seg_preview_assets import ComfySegPreviewAssetPublisher def test_publisher_stores_both_assets_and_builds_versioned_manifest( diff --git a/tests/test_seg_visualization.py b/tests/media/test_seg_visualization.py similarity index 100% rename from tests/test_seg_visualization.py rename to tests/media/test_seg_visualization.py diff --git a/tests/test_simple_preview_segs_node.py b/tests/media/test_simple_preview_segs_node.py similarity index 97% rename from tests/test_simple_preview_segs_node.py rename to tests/media/test_simple_preview_segs_node.py index 4d4132c..239bb09 100644 --- a/tests/test_simple_preview_segs_node.py +++ b/tests/media/test_simple_preview_segs_node.py @@ -11,9 +11,9 @@ from typing import ClassVar import torch from pytest import MonkeyPatch +from simple_syrup.domain.seg_preview import SegPreviewDocument from simple_syrup.nodes.simple_preview_segs import SimplePreviewSEGS from simple_syrup.runtime.seg_preview_assets import SegPreviewPublication -from simple_syrup.services.simple_preview_segs_service import SegPreviewDocument def test_node_declares_terminal_passthrough_preview_contract() -> None: diff --git a/tests/test_simple_preview_segs_service.py b/tests/media/test_simple_preview_segs_service.py similarity index 100% rename from tests/test_simple_preview_segs_service.py rename to tests/media/test_simple_preview_segs_service.py diff --git a/tests/test_simple_vae_encode_node.py b/tests/media/test_simple_vae_encode_node.py similarity index 100% rename from tests/test_simple_vae_encode_node.py rename to tests/media/test_simple_vae_encode_node.py diff --git a/tests/test_upscale_latent_from_image_node.py b/tests/media/test_upscale_latent_from_image_node.py similarity index 100% rename from tests/test_upscale_latent_from_image_node.py rename to tests/media/test_upscale_latent_from_image_node.py diff --git a/tests/models/__init__.py b/tests/models/__init__.py new file mode 100644 index 0000000..01fbdd1 --- /dev/null +++ b/tests/models/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own models test behavior.""" diff --git a/tests/models/loading/__init__.py b/tests/models/loading/__init__.py new file mode 100644 index 0000000..5448497 --- /dev/null +++ b/tests/models/loading/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own models loading test behavior.""" diff --git a/tests/test_auto_model_cache.py b/tests/models/loading/test_auto_model_cache.py similarity index 98% rename from tests/test_auto_model_cache.py rename to tests/models/loading/test_auto_model_cache.py index 2e4ef1e..5c7764a 100644 --- a/tests/test_auto_model_cache.py +++ b/tests/models/loading/test_auto_model_cache.py @@ -10,13 +10,13 @@ import json from pathlib import Path import pytest +from support.helpers import FakeFolderPaths from simple_syrup.runtime.auto_model_cache import ( AutoModelCache, AutoModelCacheEntry, AutoModelCacheError, ) -from test_helpers import FakeFolderPaths def test_cache_uses_comfy_user_directory(tmp_path: Path) -> None: diff --git a/tests/test_auto_model_choices.py b/tests/models/loading/test_auto_model_choices.py similarity index 100% rename from tests/test_auto_model_choices.py rename to tests/models/loading/test_auto_model_choices.py diff --git a/tests/test_auto_model_resolver.py b/tests/models/loading/test_auto_model_resolver.py similarity index 99% rename from tests/test_auto_model_resolver.py rename to tests/models/loading/test_auto_model_resolver.py index 2fa0a03..8b49d9a 100644 --- a/tests/test_auto_model_resolver.py +++ b/tests/models/loading/test_auto_model_resolver.py @@ -10,6 +10,7 @@ import hashlib from pathlib import Path import pytest +from support.helpers import FakeFolderPaths from simple_syrup.runtime.auto_model_artifact import AutoModelArtifact from simple_syrup.runtime.auto_model_cache import AutoModelCache, AutoModelCacheEntry @@ -21,7 +22,6 @@ from simple_syrup.runtime.auto_model_resolver import ( relative_model_name, ) from simple_syrup.runtime.model_downloads import DownloadRequest, DownloadResult -from test_helpers import FakeFolderPaths class RecordingDownloader: diff --git a/tests/test_bert_resolver.py b/tests/models/loading/test_bert_resolver.py similarity index 98% rename from tests/test_bert_resolver.py rename to tests/models/loading/test_bert_resolver.py index bffd883..dd6f971 100644 --- a/tests/test_bert_resolver.py +++ b/tests/models/loading/test_bert_resolver.py @@ -9,10 +9,10 @@ from __future__ import annotations from pathlib import Path import pytest +from support.helpers import FakeFolderPaths from simple_syrup.runtime.bert_resolver import BertResolver, is_valid_bert_directory from simple_syrup.runtime.model_downloads import DownloadRequest, DownloadResult -from test_helpers import FakeFolderPaths class RecordingDownloader: diff --git a/tests/test_checkpoint_loader.py b/tests/models/loading/test_checkpoint_loader.py similarity index 100% rename from tests/test_checkpoint_loader.py rename to tests/models/loading/test_checkpoint_loader.py diff --git a/tests/test_checkpoint_quantizer.py b/tests/models/loading/test_checkpoint_quantizer.py similarity index 100% rename from tests/test_checkpoint_quantizer.py rename to tests/models/loading/test_checkpoint_quantizer.py diff --git a/tests/test_clip_type_support.py b/tests/models/loading/test_clip_type_support.py similarity index 100% rename from tests/test_clip_type_support.py rename to tests/models/loading/test_clip_type_support.py diff --git a/tests/test_diffusion_model_metadata.py b/tests/models/loading/test_diffusion_model_metadata.py similarity index 100% rename from tests/test_diffusion_model_metadata.py rename to tests/models/loading/test_diffusion_model_metadata.py diff --git a/tests/test_flux_artifacts.py b/tests/models/loading/test_flux_artifacts.py similarity index 100% rename from tests/test_flux_artifacts.py rename to tests/models/loading/test_flux_artifacts.py diff --git a/tests/test_flux_loader_services.py b/tests/models/loading/test_flux_loader_services.py similarity index 100% rename from tests/test_flux_loader_services.py rename to tests/models/loading/test_flux_loader_services.py diff --git a/tests/test_flux_model_inspector.py b/tests/models/loading/test_flux_model_inspector.py similarity index 100% rename from tests/test_flux_model_inspector.py rename to tests/models/loading/test_flux_model_inspector.py diff --git a/tests/test_flux_profiles.py b/tests/models/loading/test_flux_profiles.py similarity index 100% rename from tests/test_flux_profiles.py rename to tests/models/loading/test_flux_profiles.py diff --git a/tests/test_flux_runtime_loaders.py b/tests/models/loading/test_flux_runtime_loaders.py similarity index 100% rename from tests/test_flux_runtime_loaders.py rename to tests/models/loading/test_flux_runtime_loaders.py diff --git a/tests/test_krea2_artifacts.py b/tests/models/loading/test_krea2_artifacts.py similarity index 100% rename from tests/test_krea2_artifacts.py rename to tests/models/loading/test_krea2_artifacts.py diff --git a/tests/test_krea2_loader_service.py b/tests/models/loading/test_krea2_loader_service.py similarity index 100% rename from tests/test_krea2_loader_service.py rename to tests/models/loading/test_krea2_loader_service.py diff --git a/tests/test_loaded_models.py b/tests/models/loading/test_loaded_models.py similarity index 100% rename from tests/test_loaded_models.py rename to tests/models/loading/test_loaded_models.py diff --git a/tests/test_model_catalog.py b/tests/models/loading/test_model_catalog.py similarity index 100% rename from tests/test_model_catalog.py rename to tests/models/loading/test_model_catalog.py diff --git a/tests/test_model_choices.py b/tests/models/loading/test_model_choices.py similarity index 100% rename from tests/test_model_choices.py rename to tests/models/loading/test_model_choices.py diff --git a/tests/test_model_device_manager.py b/tests/models/loading/test_model_device_manager.py similarity index 100% rename from tests/test_model_device_manager.py rename to tests/models/loading/test_model_device_manager.py diff --git a/tests/test_model_downloads.py b/tests/models/loading/test_model_downloads.py similarity index 100% rename from tests/test_model_downloads.py rename to tests/models/loading/test_model_downloads.py diff --git a/tests/test_model_family_profile.py b/tests/models/loading/test_model_family_profile.py similarity index 100% rename from tests/test_model_family_profile.py rename to tests/models/loading/test_model_family_profile.py diff --git a/tests/test_model_folders.py b/tests/models/loading/test_model_folders.py similarity index 98% rename from tests/test_model_folders.py rename to tests/models/loading/test_model_folders.py index 79657bd..827d9fc 100644 --- a/tests/test_model_folders.py +++ b/tests/models/loading/test_model_folders.py @@ -8,13 +8,14 @@ from __future__ import annotations from pathlib import Path +from support.helpers import FakeFolderPaths + from simple_syrup.runtime.model_folders import ( get_model_folder_paths, nonrecursive_model_files, register_required_model_folders, resolve_model_file, ) -from test_helpers import FakeFolderPaths def test_register_required_model_folders_adds_missing_folders(tmp_path: Path) -> None: diff --git a/tests/test_model_instance_cache.py b/tests/models/loading/test_model_instance_cache.py similarity index 100% rename from tests/test_model_instance_cache.py rename to tests/models/loading/test_model_instance_cache.py diff --git a/tests/test_model_metadata.py b/tests/models/loading/test_model_metadata.py similarity index 96% rename from tests/test_model_metadata.py rename to tests/models/loading/test_model_metadata.py index 5b7b4cc..3004568 100644 --- a/tests/test_model_metadata.py +++ b/tests/models/loading/test_model_metadata.py @@ -9,8 +9,9 @@ from __future__ import annotations import json from pathlib import Path +from support.helpers import FakeFolderPaths + from simple_syrup.runtime.model_metadata import GroundedSAMModelMetadata -from test_helpers import FakeFolderPaths def test_model_metadata_reports_source_and_expected_paths(tmp_path: Path) -> None: diff --git a/tests/test_optional_triton_runtime.py b/tests/models/loading/test_optional_triton_runtime.py similarity index 98% rename from tests/test_optional_triton_runtime.py rename to tests/models/loading/test_optional_triton_runtime.py index fd51d72..f031cfa 100644 --- a/tests/test_optional_triton_runtime.py +++ b/tests/models/loading/test_optional_triton_runtime.py @@ -8,9 +8,9 @@ from __future__ import annotations import subprocess import sys -from pathlib import Path import pytest +from support.repository import REPOSITORY_ROOT from simple_syrup.runtime.regional_lora.triton_runtime import ( TritonRuntimeResolver, @@ -52,7 +52,7 @@ assert not any(name == "triton" or name.startswith("triton.") for name in sys.mo """ completed = subprocess.run( [sys.executable, "-c", script], - cwd=Path(__file__).resolve().parents[1], + cwd=REPOSITORY_ROOT, capture_output=True, text=True, timeout=60, diff --git a/tests/test_quantization_capabilities.py b/tests/models/loading/test_quantization_capabilities.py similarity index 100% rename from tests/test_quantization_capabilities.py rename to tests/models/loading/test_quantization_capabilities.py diff --git a/tests/test_quantized_model_resolver.py b/tests/models/loading/test_quantized_model_resolver.py similarity index 92% rename from tests/test_quantized_model_resolver.py rename to tests/models/loading/test_quantized_model_resolver.py index 9003146..35bf582 100644 --- a/tests/test_quantized_model_resolver.py +++ b/tests/models/loading/test_quantized_model_resolver.py @@ -8,7 +8,6 @@ from __future__ import annotations import logging import threading -import time from dataclasses import dataclass, field from pathlib import Path @@ -63,10 +62,17 @@ class FakeQuantizationResult: class RecordingQuantizer: """Write a tiny artifact and record conversion calls.""" - def __init__(self, delay: float = 0.0, fail: bool = False) -> None: - """Create a fake with optional concurrency delay or failure.""" + def __init__( + self, + *, + entered: threading.Event | None = None, + release: threading.Event | None = None, + fail: bool = False, + ) -> None: + """Create a fake with optional concurrency coordination or failure.""" - self.delay = delay + self.entered = entered + self.release = release self.fail = fail self.calls: list[SourceCheckpointIdentity] = [] self._lock = threading.Lock() @@ -87,8 +93,10 @@ class RecordingQuantizer: del profile, recipe, progress, progress_base, progress_total with self._lock: self.calls.append(source) - if self.delay: - time.sleep(self.delay) + if self.entered is not None: + self.entered.set() + if self.release is not None and not self.release.wait(timeout=5.0): + raise TimeoutError("test did not release the recording quantizer") if self.fail: raise RuntimeError("conversion failed") destination_path.write_bytes(b"quantized") @@ -316,7 +324,12 @@ def test_concurrent_requests_share_one_generated_artifact(tmp_path: Path) -> Non source = tmp_path / "anima.safetensors" source.write_bytes(b"same-source") repository = QuantCacheRepository(tmp_path / "SyrupQuants") - quantizer = RecordingQuantizer(delay=0.2) + quantizer_entered = threading.Event() + release_quantizer = threading.Event() + quantizer = RecordingQuantizer( + entered=quantizer_entered, + release=release_quantizer, + ) leases = QuantCacheLeaseRegistry() resolver = QuantizedModelResolver( repository=repository, @@ -349,8 +362,10 @@ def test_concurrent_requests_share_one_generated_artifact(tmp_path: Path) -> Non failures.append(error) threads = [threading.Thread(target=resolve_once) for _ in range(2)] - for thread in threads: - thread.start() + threads[0].start() + assert quantizer_entered.wait(timeout=5.0) + threads[1].start() + release_quantizer.set() for thread in threads: thread.join(timeout=5) diff --git a/tests/test_simple_load_anima_node.py b/tests/models/loading/test_simple_load_anima_node.py similarity index 100% rename from tests/test_simple_load_anima_node.py rename to tests/models/loading/test_simple_load_anima_node.py diff --git a/tests/test_simple_load_checkpoint_node.py b/tests/models/loading/test_simple_load_checkpoint_node.py similarity index 100% rename from tests/test_simple_load_checkpoint_node.py rename to tests/models/loading/test_simple_load_checkpoint_node.py diff --git a/tests/test_simple_load_checkpoint_v3_node.py b/tests/models/loading/test_simple_load_checkpoint_v3_node.py similarity index 100% rename from tests/test_simple_load_checkpoint_v3_node.py rename to tests/models/loading/test_simple_load_checkpoint_v3_node.py diff --git a/tests/test_simple_load_flux_nodes.py b/tests/models/loading/test_simple_load_flux_nodes.py similarity index 100% rename from tests/test_simple_load_flux_nodes.py rename to tests/models/loading/test_simple_load_flux_nodes.py diff --git a/tests/test_simple_load_krea2_node.py b/tests/models/loading/test_simple_load_krea2_node.py similarity index 100% rename from tests/test_simple_load_krea2_node.py rename to tests/models/loading/test_simple_load_krea2_node.py diff --git a/tests/test_vae_loader.py b/tests/models/loading/test_vae_loader.py similarity index 100% rename from tests/test_vae_loader.py rename to tests/models/loading/test_vae_loader.py diff --git a/tests/models/patching/__init__.py b/tests/models/patching/__init__.py new file mode 100644 index 0000000..e53c2bd --- /dev/null +++ b/tests/models/patching/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own models patching test behavior.""" diff --git a/tests/test_clip_patcher_model_alignment.py b/tests/models/patching/test_clip_patcher_model_alignment.py similarity index 100% rename from tests/test_clip_patcher_model_alignment.py rename to tests/models/patching/test_clip_patcher_model_alignment.py diff --git a/tests/test_clip_patcher_mutations.py b/tests/models/patching/test_clip_patcher_mutations.py similarity index 100% rename from tests/test_clip_patcher_mutations.py rename to tests/models/patching/test_clip_patcher_mutations.py diff --git a/tests/test_comfy_patcher_lifecycle.py b/tests/models/patching/test_comfy_patcher_lifecycle.py similarity index 100% rename from tests/test_comfy_patcher_lifecycle.py rename to tests/models/patching/test_comfy_patcher_lifecycle.py diff --git a/tests/test_lora_execution_probe.py b/tests/models/patching/test_lora_execution_probe.py similarity index 100% rename from tests/test_lora_execution_probe.py rename to tests/models/patching/test_lora_execution_probe.py diff --git a/tests/test_lora_patch_snapshot.py b/tests/models/patching/test_lora_patch_snapshot.py similarity index 100% rename from tests/test_lora_patch_snapshot.py rename to tests/models/patching/test_lora_patch_snapshot.py diff --git a/tests/test_model_attn1_replacement_mutation.py b/tests/models/patching/test_model_attn1_replacement_mutation.py similarity index 100% rename from tests/test_model_attn1_replacement_mutation.py rename to tests/models/patching/test_model_attn1_replacement_mutation.py diff --git a/tests/test_model_keyed_callback_mutation.py b/tests/models/patching/test_model_keyed_callback_mutation.py similarity index 100% rename from tests/test_model_keyed_callback_mutation.py rename to tests/models/patching/test_model_keyed_callback_mutation.py diff --git a/tests/test_model_modifier_snapshot_probe.py b/tests/models/patching/test_model_modifier_snapshot_probe.py similarity index 100% rename from tests/test_model_modifier_snapshot_probe.py rename to tests/models/patching/test_model_modifier_snapshot_probe.py diff --git a/tests/test_model_object_patch_batch.py b/tests/models/patching/test_model_object_patch_batch.py similarity index 100% rename from tests/test_model_object_patch_batch.py rename to tests/models/patching/test_model_object_patch_batch.py diff --git a/tests/models/patching/test_model_object_patch_mutations.py b/tests/models/patching/test_model_object_patch_mutations.py new file mode 100644 index 0000000..4a1f1ce --- /dev/null +++ b/tests/models/patching/test_model_object_patch_mutations.py @@ -0,0 +1,167 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Characterize concrete MODEL patcher mutation behavior.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.runtime.model_patcher_mutations import ( + ModelExactObjectPatchMutation, + ModelSharedObjectPatchMutation, +) + + +@pytest.mark.parametrize("path", ["", ".weight", "weight.", "model..weight"]) +def test_exact_object_patch_rejects_invalid_dotted_paths(path: str) -> None: + """Reject empty path segments before inspecting or changing patcher state.""" + + model = _patcher(torch.nn.Linear(1, 1)) + + with pytest.raises(ValueError, match="non-empty dotted path"): + ModelExactObjectPatchMutation(path, object(), object()).apply(model) + + assert model.object_patches == {} + + +@pytest.mark.parametrize("state_name", ["object_patches", "object_patches_backup"]) +def test_exact_object_patch_rejects_existing_active_or_backup_state( + state_name: str, +) -> None: + """Reject both forms of an already-owned exact object path.""" + + model = _patcher(torch.nn.Linear(1, 1)) + expected = model.get_model_object("weight") + getattr(model, state_name)["weight"] = expected + + with pytest.raises(ValueError, match="already has a patch"): + ModelExactObjectPatchMutation("weight", expected, object()).apply(model) + + assert ( + "weight" not in model.object_patches + or model.object_patches["weight"] is expected + ) + + +def test_exact_object_patch_rejects_changed_current_identity() -> None: + """Require the caller's exact expected object before claiming the path.""" + + model = _patcher(torch.nn.Linear(1, 1)) + + with pytest.raises(ValueError, match="does not match the expected object"): + ModelExactObjectPatchMutation("weight", object(), object()).apply(model) + + assert model.object_patches == {} + + +def test_exact_object_patch_rejects_changed_adder_signature_before_lookup() -> None: + """Validate the whole object-patch API before looking up or replacing an object.""" + + class ChangedSurface: + """Expose an incompatible object-patch adder.""" + + def __init__(self) -> None: + """Initialize valid state and untouched lookup state.""" + + self.object_patches: dict[str, object] = {} + self.object_patches_backup: dict[str, object] = {} + self.looked_up = False + + def get_model_object(self, name: str) -> object: + """Record a lookup that must never occur.""" + + del name + self.looked_up = True + return object() + + def add_object_patch(self, path: str, value: object) -> None: + """Expose deliberately changed parameter names.""" + + del path, value + + model = ChangedSurface() + + with pytest.raises(TypeError, match="unsupported signature"): + ModelExactObjectPatchMutation("weight", object(), object()).apply(model) + + assert model.looked_up is False + assert model.object_patches == {} + + +@pytest.mark.parametrize("attribute_name", ["object_patches", "object_patches_backup"]) +def test_exact_object_patch_rejects_malformed_state_dictionary( + attribute_name: str, +) -> None: + """Require both exact-path collision stores to remain dictionaries.""" + + model = _patcher(torch.nn.Linear(1, 1)) + expected = model.get_model_object("weight") + setattr(model, attribute_name, None) + + with pytest.raises(TypeError, match=f"{attribute_name} must be a dictionary"): + ModelExactObjectPatchMutation("weight", expected, object()).apply(model) + + setattr(model, attribute_name, {}) + + +def test_shared_object_patch_accepts_exact_backup_and_live_replacement() -> None: + """Register the already-live replacement when both identities remain exact.""" + + model = _patcher(torch.nn.Linear(1, 1)) + expected_backup = model.model.weight + replacement = torch.nn.Parameter(torch.zeros_like(expected_backup)) + model.object_patches_backup["weight"] = expected_backup + model.model.weight = replacement + + ModelSharedObjectPatchMutation( + "weight", + expected_backup, + replacement, + ).apply(model) + + assert model.object_patches["weight"] is replacement + + +@pytest.mark.parametrize("foreign_state", ["backup", "live"]) +def test_shared_object_patch_rejects_foreign_shared_state( + foreign_state: str, +) -> None: + """Fail closed when either shared identity no longer belongs to the caller.""" + + model = _patcher(torch.nn.Linear(1, 1)) + expected_backup = model.model.weight + replacement = torch.nn.Parameter(torch.zeros_like(expected_backup)) + model.object_patches_backup["weight"] = ( + object() if foreign_state == "backup" else expected_backup + ) + model.model.weight = ( + torch.nn.Parameter(torch.ones_like(expected_backup)) + if foreign_state == "live" + else replacement + ) + + expected_error = ( + "foreign shared backup" if foreign_state == "backup" else "foreign live" + ) + with pytest.raises(ValueError, match=expected_error): + ModelSharedObjectPatchMutation( + "weight", + expected_backup, + replacement, + ).apply(model) + + assert model.object_patches == {} + + +def _patcher(model: torch.nn.Module) -> Any: + """Create a real CPU Comfy MODEL patcher.""" + + from comfy.model_patcher import ModelPatcher + + device = torch.device("cpu") + return ModelPatcher(model, load_device=device, offload_device=device) diff --git a/tests/test_model_patcher_mutations.py b/tests/models/patching/test_model_patcher_mutations.py similarity index 78% rename from tests/test_model_patcher_mutations.py rename to tests/models/patching/test_model_patcher_mutations.py index 7d9df8d..1cef067 100644 --- a/tests/test_model_patcher_mutations.py +++ b/tests/models/patching/test_model_patcher_mutations.py @@ -21,7 +21,6 @@ from simple_syrup.runtime.model_patcher_mutations import ( ModelDiffusionWrapperMutation, ModelExactObjectPatchMutation, ModelKeyedWrapperMutation, - ModelSharedObjectPatchMutation, ModelUnetWrapperMutation, ) from simple_syrup.runtime.patcher_lifecycle import PATCHER_LIFECYCLE @@ -557,147 +556,6 @@ def test_attn2_mutation_rejects_changed_output_signature_before_input_call() -> assert model.calls == [] -@pytest.mark.parametrize("path", ["", ".weight", "weight.", "model..weight"]) -def test_exact_object_patch_rejects_invalid_dotted_paths(path: str) -> None: - """Reject empty path segments before inspecting or changing patcher state.""" - - model = _patcher(torch.nn.Linear(1, 1)) - - with pytest.raises(ValueError, match="non-empty dotted path"): - ModelExactObjectPatchMutation(path, object(), object()).apply(model) - - assert model.object_patches == {} - - -@pytest.mark.parametrize("state_name", ["object_patches", "object_patches_backup"]) -def test_exact_object_patch_rejects_existing_active_or_backup_state( - state_name: str, -) -> None: - """Reject both forms of an already-owned exact object path.""" - - model = _patcher(torch.nn.Linear(1, 1)) - expected = model.get_model_object("weight") - getattr(model, state_name)["weight"] = expected - - with pytest.raises(ValueError, match="already has a patch"): - ModelExactObjectPatchMutation("weight", expected, object()).apply(model) - - assert ( - "weight" not in model.object_patches - or model.object_patches["weight"] is expected - ) - - -def test_exact_object_patch_rejects_changed_current_identity() -> None: - """Require the caller's exact expected object before claiming the path.""" - - model = _patcher(torch.nn.Linear(1, 1)) - - with pytest.raises(ValueError, match="does not match the expected object"): - ModelExactObjectPatchMutation("weight", object(), object()).apply(model) - - assert model.object_patches == {} - - -def test_exact_object_patch_rejects_changed_adder_signature_before_lookup() -> None: - """Validate the whole object-patch API before looking up or replacing an object.""" - - class ChangedSurface: - """Expose an incompatible object-patch adder.""" - - def __init__(self) -> None: - """Initialize valid state and untouched lookup state.""" - - self.object_patches: dict[str, object] = {} - self.object_patches_backup: dict[str, object] = {} - self.looked_up = False - - def get_model_object(self, name: str) -> object: - """Record a lookup that must never occur.""" - - del name - self.looked_up = True - return object() - - def add_object_patch(self, path: str, value: object) -> None: - """Expose deliberately changed parameter names.""" - - del path, value - - model = ChangedSurface() - - with pytest.raises(TypeError, match="unsupported signature"): - ModelExactObjectPatchMutation("weight", object(), object()).apply(model) - - assert model.looked_up is False - assert model.object_patches == {} - - -@pytest.mark.parametrize("attribute_name", ["object_patches", "object_patches_backup"]) -def test_exact_object_patch_rejects_malformed_state_dictionary( - attribute_name: str, -) -> None: - """Require both exact-path collision stores to remain dictionaries.""" - - model = _patcher(torch.nn.Linear(1, 1)) - expected = model.get_model_object("weight") - setattr(model, attribute_name, None) - - with pytest.raises(TypeError, match=f"{attribute_name} must be a dictionary"): - ModelExactObjectPatchMutation("weight", expected, object()).apply(model) - - setattr(model, attribute_name, {}) - - -def test_shared_object_patch_accepts_exact_backup_and_live_replacement() -> None: - """Register the already-live replacement when both identities remain exact.""" - - model = _patcher(torch.nn.Linear(1, 1)) - expected_backup = model.model.weight - replacement = torch.nn.Parameter(torch.zeros_like(expected_backup)) - model.object_patches_backup["weight"] = expected_backup - model.model.weight = replacement - - ModelSharedObjectPatchMutation( - "weight", - expected_backup, - replacement, - ).apply(model) - - assert model.object_patches["weight"] is replacement - - -@pytest.mark.parametrize("foreign_state", ["backup", "live"]) -def test_shared_object_patch_rejects_foreign_shared_state( - foreign_state: str, -) -> None: - """Fail closed when either shared identity no longer belongs to the caller.""" - - model = _patcher(torch.nn.Linear(1, 1)) - expected_backup = model.model.weight - replacement = torch.nn.Parameter(torch.zeros_like(expected_backup)) - model.object_patches_backup["weight"] = ( - object() if foreign_state == "backup" else expected_backup - ) - model.model.weight = ( - torch.nn.Parameter(torch.ones_like(expected_backup)) - if foreign_state == "live" - else replacement - ) - - expected_error = ( - "foreign shared backup" if foreign_state == "backup" else "foreign live" - ) - with pytest.raises(ValueError, match=expected_error): - ModelSharedObjectPatchMutation( - "weight", - expected_backup, - replacement, - ).apply(model) - - assert model.object_patches == {} - - def _patcher(model: torch.nn.Module) -> Any: """Create a real CPU Comfy MODEL patcher.""" diff --git a/tests/test_patcher_lifecycle_policy.py b/tests/models/patching/test_patcher_lifecycle_policy.py similarity index 98% rename from tests/test_patcher_lifecycle_policy.py rename to tests/models/patching/test_patcher_lifecycle_policy.py index 7c96476..e6c4812 100644 --- a/tests/test_patcher_lifecycle_policy.py +++ b/tests/models/patching/test_patcher_lifecycle_policy.py @@ -8,9 +8,10 @@ from __future__ import annotations import ast from collections import Counter -from pathlib import Path -PROJECT_ROOT = Path(__file__).resolve().parents[1] +from support.repository import REPOSITORY_ROOT + +PROJECT_ROOT = REPOSITORY_ROOT SOURCE_ROOT = PROJECT_ROOT / "simple_syrup" LIFECYCLE_MODULE = "simple_syrup/runtime/patcher_lifecycle.py" MUTATION_MODULES = frozenset( diff --git a/tests/models/quantization/__init__.py b/tests/models/quantization/__init__.py new file mode 100644 index 0000000..18d0179 --- /dev/null +++ b/tests/models/quantization/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own models quantization test behavior.""" diff --git a/tests/test_comfy_safetensors_dtypes.py b/tests/models/quantization/test_comfy_safetensors_dtypes.py similarity index 100% rename from tests/test_comfy_safetensors_dtypes.py rename to tests/models/quantization/test_comfy_safetensors_dtypes.py diff --git a/tests/test_model_quantization.py b/tests/models/quantization/test_model_quantization.py similarity index 100% rename from tests/test_model_quantization.py rename to tests/models/quantization/test_model_quantization.py diff --git a/tests/test_quant_cache.py b/tests/models/quantization/test_quant_cache.py similarity index 100% rename from tests/test_quant_cache.py rename to tests/models/quantization/test_quant_cache.py diff --git a/tests/node_api/__init__.py b/tests/node_api/__init__.py new file mode 100644 index 0000000..cb91570 --- /dev/null +++ b/tests/node_api/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own nodes test behavior.""" diff --git a/tests/test_batch_region_conditioning_node.py b/tests/node_api/test_batch_region_conditioning_node.py similarity index 100% rename from tests/test_batch_region_conditioning_node.py rename to tests/node_api/test_batch_region_conditioning_node.py diff --git a/tests/test_batch_region_conditioning_v3_node.py b/tests/node_api/test_batch_region_conditioning_v3_node.py similarity index 100% rename from tests/test_batch_region_conditioning_v3_node.py rename to tests/node_api/test_batch_region_conditioning_v3_node.py diff --git a/tests/test_compose_regional_conditioning_v3_node.py b/tests/node_api/test_compose_regional_conditioning_v3_node.py similarity index 100% rename from tests/test_compose_regional_conditioning_v3_node.py rename to tests/node_api/test_compose_regional_conditioning_v3_node.py diff --git a/tests/test_latent_diagnostics_node.py b/tests/node_api/test_latent_diagnostics_node.py similarity index 100% rename from tests/test_latent_diagnostics_node.py rename to tests/node_api/test_latent_diagnostics_node.py diff --git a/tests/test_latent_diagnostics_service.py b/tests/node_api/test_latent_diagnostics_service.py similarity index 100% rename from tests/test_latent_diagnostics_service.py rename to tests/node_api/test_latent_diagnostics_service.py diff --git a/tests/test_legacy_node_v3_adapter.py b/tests/node_api/test_legacy_node_v3_adapter.py similarity index 100% rename from tests/test_legacy_node_v3_adapter.py rename to tests/node_api/test_legacy_node_v3_adapter.py diff --git a/tests/test_seed_node.py b/tests/node_api/test_seed_node.py similarity index 100% rename from tests/test_seed_node.py rename to tests/node_api/test_seed_node.py diff --git a/tests/test_seed_variation_domain.py b/tests/node_api/test_seed_variation_domain.py similarity index 100% rename from tests/test_seed_variation_domain.py rename to tests/node_api/test_seed_variation_domain.py diff --git a/tests/test_seed_variation_runtime.py b/tests/node_api/test_seed_variation_runtime.py similarity index 100% rename from tests/test_seed_variation_runtime.py rename to tests/node_api/test_seed_variation_runtime.py diff --git a/tests/test_seed_variation_v3_node.py b/tests/node_api/test_seed_variation_v3_node.py similarity index 100% rename from tests/test_seed_variation_v3_node.py rename to tests/node_api/test_seed_variation_v3_node.py diff --git a/tests/test_vae_decode_options_v3_node.py b/tests/node_api/test_vae_decode_options_v3_node.py similarity index 100% rename from tests/test_vae_decode_options_v3_node.py rename to tests/node_api/test_vae_decode_options_v3_node.py diff --git a/tests/test_vae_encode_options_v3_node.py b/tests/node_api/test_vae_encode_options_v3_node.py similarity index 100% rename from tests/test_vae_encode_options_v3_node.py rename to tests/node_api/test_vae_encode_options_v3_node.py diff --git a/tests/test_vae_options_node.py b/tests/node_api/test_vae_options_node.py similarity index 100% rename from tests/test_vae_options_node.py rename to tests/node_api/test_vae_options_node.py diff --git a/tests/prompting/__init__.py b/tests/prompting/__init__.py new file mode 100644 index 0000000..8c01f2d --- /dev/null +++ b/tests/prompting/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own prompting test behavior.""" diff --git a/tests/prompting/conditioning/__init__.py b/tests/prompting/conditioning/__init__.py new file mode 100644 index 0000000..a3f1ca7 --- /dev/null +++ b/tests/prompting/conditioning/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own prompting conditioning test behavior.""" diff --git a/tests/test_conditioning_batch.py b/tests/prompting/conditioning/test_conditioning_batch.py similarity index 100% rename from tests/test_conditioning_batch.py rename to tests/prompting/conditioning/test_conditioning_batch.py diff --git a/tests/test_conditioning_batch_bridge.py b/tests/prompting/conditioning/test_conditioning_batch_bridge.py similarity index 100% rename from tests/test_conditioning_batch_bridge.py rename to tests/prompting/conditioning/test_conditioning_batch_bridge.py diff --git a/tests/test_conditioning_batch_pack_node.py b/tests/prompting/conditioning/test_conditioning_batch_pack_node.py similarity index 100% rename from tests/test_conditioning_batch_pack_node.py rename to tests/prompting/conditioning/test_conditioning_batch_pack_node.py diff --git a/tests/test_conditioning_batch_snapshot_probe.py b/tests/prompting/conditioning/test_conditioning_batch_snapshot_probe.py similarity index 100% rename from tests/test_conditioning_batch_snapshot_probe.py rename to tests/prompting/conditioning/test_conditioning_batch_snapshot_probe.py diff --git a/tests/test_conditioning_schedule.py b/tests/prompting/conditioning/test_conditioning_schedule.py similarity index 100% rename from tests/test_conditioning_schedule.py rename to tests/prompting/conditioning/test_conditioning_schedule.py diff --git a/tests/test_conditioning_schedule_comfy_equivalence.py b/tests/prompting/conditioning/test_conditioning_schedule_comfy_equivalence.py similarity index 100% rename from tests/test_conditioning_schedule_comfy_equivalence.py rename to tests/prompting/conditioning/test_conditioning_schedule_comfy_equivalence.py diff --git a/tests/test_global_context_schedule.py b/tests/prompting/conditioning/test_global_context_schedule.py similarity index 100% rename from tests/test_global_context_schedule.py rename to tests/prompting/conditioning/test_global_context_schedule.py diff --git a/tests/prompting/negpip/__init__.py b/tests/prompting/negpip/__init__.py new file mode 100644 index 0000000..ae7df09 --- /dev/null +++ b/tests/prompting/negpip/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own prompting negpip test behavior.""" diff --git a/tests/test_clip_schedule_snapshot_probe.py b/tests/prompting/negpip/test_clip_schedule_snapshot_probe.py similarity index 100% rename from tests/test_clip_schedule_snapshot_probe.py rename to tests/prompting/negpip/test_clip_schedule_snapshot_probe.py diff --git a/tests/test_negative_prompt_weights.py b/tests/prompting/negpip/test_negative_prompt_weights.py similarity index 100% rename from tests/test_negative_prompt_weights.py rename to tests/prompting/negpip/test_negative_prompt_weights.py diff --git a/tests/test_negpip_integration_workflow.py b/tests/prompting/negpip/test_negpip_integration_workflow.py similarity index 100% rename from tests/test_negpip_integration_workflow.py rename to tests/prompting/negpip/test_negpip_integration_workflow.py diff --git a/tests/test_negpip_model_service.py b/tests/prompting/negpip/test_negpip_model_service.py similarity index 100% rename from tests/test_negpip_model_service.py rename to tests/prompting/negpip/test_negpip_model_service.py diff --git a/tests/test_negpip_runtime.py b/tests/prompting/negpip/test_negpip_runtime.py similarity index 100% rename from tests/test_negpip_runtime.py rename to tests/prompting/negpip/test_negpip_runtime.py diff --git a/tests/test_negpip_runtime_probe.py b/tests/prompting/negpip/test_negpip_runtime_probe.py similarity index 100% rename from tests/test_negpip_runtime_probe.py rename to tests/prompting/negpip/test_negpip_runtime_probe.py diff --git a/tests/test_negpip_visual_proof.py b/tests/prompting/negpip/test_negpip_visual_proof.py similarity index 100% rename from tests/test_negpip_visual_proof.py rename to tests/prompting/negpip/test_negpip_visual_proof.py diff --git a/tests/test_ppm_negpip_interop.py b/tests/prompting/negpip/test_ppm_negpip_interop.py similarity index 100% rename from tests/test_ppm_negpip_interop.py rename to tests/prompting/negpip/test_ppm_negpip_interop.py diff --git a/tests/test_scheduled_clip_conditioning_metadata.py b/tests/prompting/negpip/test_scheduled_clip_conditioning_metadata.py similarity index 100% rename from tests/test_scheduled_clip_conditioning_metadata.py rename to tests/prompting/negpip/test_scheduled_clip_conditioning_metadata.py diff --git a/tests/prompting/prompt_control/__init__.py b/tests/prompting/prompt_control/__init__.py new file mode 100644 index 0000000..f3ea796 --- /dev/null +++ b/tests/prompting/prompt_control/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own prompting prompt control test behavior.""" diff --git a/tests/prompting/prompt_control/support/__init__.py b/tests/prompting/prompt_control/support/__init__.py new file mode 100644 index 0000000..b8ffd3b --- /dev/null +++ b/tests/prompting/prompt_control/support/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own prompting prompt control support test behavior.""" diff --git a/tests/test_prompt_control_evidence_validation.py b/tests/prompting/prompt_control/support/evidence.py similarity index 70% rename from tests/test_prompt_control_evidence_validation.py rename to tests/prompting/prompt_control/support/evidence.py index 7e3fd32..8fcee78 100644 --- a/tests/test_prompt_control_evidence_validation.py +++ b/tests/prompting/prompt_control/support/evidence.py @@ -2,52 +2,17 @@ # Copyright (C) 2026 Artificial Sweetener and contributors # SPDX-License-Identifier: AGPL-3.0-or-later -"""Verify exact Prompt Control characterization acceptance policy.""" +"""Build shared Prompt Control characterization evidence for tests.""" from __future__ import annotations import uuid -import pytest - -from tools.prompt_control_characterization.cases import PromptControlCase, cases -from tools.prompt_control_characterization.evidence_validation import ( - validate_evidence, -) +from tools.prompt_control_characterization.cases import PromptControlCase from tools.prompt_control_characterization.history_outputs import PromptControlOutputs -def test_static_text_evidence_accepts_exact_entries_and_stable_uuids() -> None: - """Accept the base schedule while proving UUIDv4 entry identity.""" - - case = _case("text-static") - outputs = _outputs(case) - validate_evidence(case, outputs) - - -def test_evidence_rejects_conditioning_boundary_drift() -> None: - """Fail when Prompt Control changes an exact range endpoint.""" - - case = _case("text-static") - outputs = _outputs(case) - positive = outputs.snapshot["positive"] - assert isinstance(positive, list) - first = positive[0] - assert isinstance(first, dict) - metadata = first["metadata"] - assert isinstance(metadata, dict) - metadata["end_percent"] = 0.9 - with pytest.raises(ValueError, match="end_percent"): - validate_evidence(case, outputs) - - -def _case(case_id: str) -> PromptControlCase: - """Return one case by stable identity.""" - - return next(case for case in cases() if case.case_id == case_id) - - -def _outputs(case: PromptControlCase) -> PromptControlOutputs: +def prompt_control_outputs(case: PromptControlCase) -> PromptControlOutputs: """Build minimal complete text-only host evidence.""" negative_uuids = [str(uuid.uuid4()) for _ in case.expected_negative] diff --git a/tests/prompt_control_attention_coupling_values.py b/tests/prompting/prompt_control/support/prompt_control_attention_coupling_values.py similarity index 100% rename from tests/prompt_control_attention_coupling_values.py rename to tests/prompting/prompt_control/support/prompt_control_attention_coupling_values.py diff --git a/tests/test_encode_prompt_batch_node.py b/tests/prompting/prompt_control/test_encode_prompt_batch_node.py similarity index 100% rename from tests/test_encode_prompt_batch_node.py rename to tests/prompting/prompt_control/test_encode_prompt_batch_node.py diff --git a/tests/test_encode_prompt_batch_with_prompt_control_node.py b/tests/prompting/prompt_control/test_encode_prompt_batch_with_prompt_control_node.py similarity index 100% rename from tests/test_encode_prompt_batch_with_prompt_control_node.py rename to tests/prompting/prompt_control/test_encode_prompt_batch_with_prompt_control_node.py diff --git a/tests/test_prompt_batch_parser.py b/tests/prompting/prompt_control/test_prompt_batch_parser.py similarity index 100% rename from tests/test_prompt_batch_parser.py rename to tests/prompting/prompt_control/test_prompt_batch_parser.py diff --git a/tests/test_prompt_composition.py b/tests/prompting/prompt_control/test_prompt_composition.py similarity index 100% rename from tests/test_prompt_composition.py rename to tests/prompting/prompt_control/test_prompt_composition.py diff --git a/tests/test_prompt_control_attention_coupling_cli.py b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_cli.py similarity index 97% rename from tests/test_prompt_control_attention_coupling_cli.py rename to tests/prompting/prompt_control/test_prompt_control_attention_coupling_cli.py index 6d22f15..19f7362 100644 --- a/tests/test_prompt_control_attention_coupling_cli.py +++ b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_cli.py @@ -21,6 +21,7 @@ def test_cli_help_exposes_managed_source_baseline_and_timeout_boundaries() -> No check=True, capture_output=True, text=True, + timeout=60.0, ) assert "--comfy-root" in completed.stdout diff --git a/tests/test_prompt_control_attention_coupling_history.py b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_history.py similarity index 100% rename from tests/test_prompt_control_attention_coupling_history.py rename to tests/prompting/prompt_control/test_prompt_control_attention_coupling_history.py diff --git a/tests/test_prompt_control_attention_coupling_matrix.py b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_matrix.py similarity index 100% rename from tests/test_prompt_control_attention_coupling_matrix.py rename to tests/prompting/prompt_control/test_prompt_control_attention_coupling_matrix.py diff --git a/tests/test_prompt_control_attention_coupling_results.py b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_results.py similarity index 97% rename from tests/test_prompt_control_attention_coupling_results.py rename to tests/prompting/prompt_control/test_prompt_control_attention_coupling_results.py index bdd93eb..2c70f83 100644 --- a/tests/test_prompt_control_attention_coupling_results.py +++ b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_results.py @@ -13,7 +13,6 @@ from typing import cast import pytest from PIL import Image -from prompt_control_attention_coupling_values import evidence_outputs from tools.comfy_api import ImageReference, JsonObject from tools.prompt_control_attention_coupling_integration.baseline import load_baseline @@ -31,6 +30,10 @@ from tools.prompt_control_characterization.source_identity import ( PromptControlSourceIdentity, ) +from .support.prompt_control_attention_coupling_values import ( + evidence_outputs, +) + pytestmark = pytest.mark.external_artifact diff --git a/tests/test_prompt_control_attention_coupling_validation.py b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_validation.py similarity index 98% rename from tests/test_prompt_control_attention_coupling_validation.py rename to tests/prompting/prompt_control/test_prompt_control_attention_coupling_validation.py index 0cd041b..ee29084 100644 --- a/tests/test_prompt_control_attention_coupling_validation.py +++ b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_validation.py @@ -9,7 +9,6 @@ from __future__ import annotations from copy import deepcopy import pytest -from prompt_control_attention_coupling_values import evidence_outputs from tools.prompt_control_attention_coupling_integration.baseline import load_baseline from tools.prompt_control_attention_coupling_integration.matrix import cases @@ -20,6 +19,10 @@ from tools.prompt_control_attention_coupling_integration.workflow import ( PromptControlAttentionWorkflowBuilder, ) +from .support.prompt_control_attention_coupling_values import ( + evidence_outputs, +) + pytestmark = pytest.mark.external_artifact diff --git a/tests/test_prompt_control_attention_coupling_workflow.py b/tests/prompting/prompt_control/test_prompt_control_attention_coupling_workflow.py similarity index 100% rename from tests/test_prompt_control_attention_coupling_workflow.py rename to tests/prompting/prompt_control/test_prompt_control_attention_coupling_workflow.py diff --git a/tests/test_prompt_control_availability.py b/tests/prompting/prompt_control/test_prompt_control_availability.py similarity index 91% rename from tests/test_prompt_control_availability.py rename to tests/prompting/prompt_control/test_prompt_control_availability.py index 876689f..6cbf00b 100644 --- a/tests/test_prompt_control_availability.py +++ b/tests/prompting/prompt_control/test_prompt_control_availability.py @@ -46,13 +46,14 @@ def test_find_prompt_control_install_reports_missing_sibling( def test_find_prompt_control_install_does_not_import_lazy_nodes( tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, ) -> None: """Availability checks inspect files without importing nodes_lazy.""" package_path = tmp_path / "comfyui-prompt-control" / "prompt_control" package_path.mkdir(parents=True) (package_path / "nodes_lazy.py").write_text("", encoding="utf-8") - sys.modules.pop("prompt_control.nodes_lazy", None) + monkeypatch.delitem(sys.modules, "prompt_control.nodes_lazy", raising=False) availability = find_prompt_control_install(custom_nodes_root=tmp_path) @@ -69,8 +70,8 @@ def test_find_prompt_control_install_detects_python_path_package( package_path = tmp_path / "prompt_control" package_path.mkdir() (package_path / "nodes_lazy.py").write_text("", encoding="utf-8") - sys.modules.pop("prompt_control", None) - sys.modules.pop("prompt_control.nodes_lazy", None) + monkeypatch.delitem(sys.modules, "prompt_control", raising=False) + monkeypatch.delitem(sys.modules, "prompt_control.nodes_lazy", raising=False) monkeypatch.syspath_prepend(str(tmp_path)) invalidate_caches() diff --git a/tests/test_prompt_control_batch_graph.py b/tests/prompting/prompt_control/test_prompt_control_batch_graph.py similarity index 99% rename from tests/test_prompt_control_batch_graph.py rename to tests/prompting/prompt_control/test_prompt_control_batch_graph.py index bcdc859..db5147b 100644 --- a/tests/test_prompt_control_batch_graph.py +++ b/tests/prompting/prompt_control/test_prompt_control_batch_graph.py @@ -13,7 +13,7 @@ from typing import Any, cast import pytest -from simple_syrup.runtime.prompt_control_batch_graph import ( +from simple_syrup.services.prompt_control_batch_graph import ( PROMPT_CONTROL_MISSING_MESSAGE, PromptControlBatchGraphBuilder, ) diff --git a/tests/test_prompt_control_characterization_cases.py b/tests/prompting/prompt_control/test_prompt_control_characterization_cases.py similarity index 100% rename from tests/test_prompt_control_characterization_cases.py rename to tests/prompting/prompt_control/test_prompt_control_characterization_cases.py diff --git a/tests/test_prompt_control_characterization_cli.py b/tests/prompting/prompt_control/test_prompt_control_characterization_cli.py similarity index 97% rename from tests/test_prompt_control_characterization_cli.py rename to tests/prompting/prompt_control/test_prompt_control_characterization_cli.py index dc8ab88..b7f7c25 100644 --- a/tests/test_prompt_control_characterization_cli.py +++ b/tests/prompting/prompt_control/test_prompt_control_characterization_cli.py @@ -16,6 +16,7 @@ def test_cli_help_exposes_explicit_external_boundaries() -> None: check=True, capture_output=True, text=True, + timeout=60.0, ) assert "--prompt-control-root" in completed.stdout assert "--server-url" in completed.stdout diff --git a/tests/test_prompt_control_characterization_workflow.py b/tests/prompting/prompt_control/test_prompt_control_characterization_workflow.py similarity index 100% rename from tests/test_prompt_control_characterization_workflow.py rename to tests/prompting/prompt_control/test_prompt_control_characterization_workflow.py diff --git a/tests/prompting/prompt_control/test_prompt_control_evidence_validation.py b/tests/prompting/prompt_control/test_prompt_control_evidence_validation.py new file mode 100644 index 0000000..4ae212e --- /dev/null +++ b/tests/prompting/prompt_control/test_prompt_control_evidence_validation.py @@ -0,0 +1,46 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Verify exact Prompt Control characterization acceptance policy.""" + +from __future__ import annotations + +import pytest + +from tools.prompt_control_characterization.cases import PromptControlCase, cases +from tools.prompt_control_characterization.evidence_validation import ( + validate_evidence, +) + +from .support.evidence import prompt_control_outputs + + +def test_static_text_evidence_accepts_exact_entries_and_stable_uuids() -> None: + """Accept the base schedule while proving UUIDv4 entry identity.""" + + case = _case("text-static") + outputs = prompt_control_outputs(case) + validate_evidence(case, outputs) + + +def test_evidence_rejects_conditioning_boundary_drift() -> None: + """Fail when Prompt Control changes an exact range endpoint.""" + + case = _case("text-static") + outputs = prompt_control_outputs(case) + positive = outputs.snapshot["positive"] + assert isinstance(positive, list) + first = positive[0] + assert isinstance(first, dict) + metadata = first["metadata"] + assert isinstance(metadata, dict) + metadata["end_percent"] = 0.9 + with pytest.raises(ValueError, match="end_percent"): + validate_evidence(case, outputs) + + +def _case(case_id: str) -> PromptControlCase: + """Return one case by stable identity.""" + + return next(case for case in cases() if case.case_id == case_id) diff --git a/tests/test_prompt_control_expansion_probe.py b/tests/prompting/prompt_control/test_prompt_control_expansion_probe.py similarity index 100% rename from tests/test_prompt_control_expansion_probe.py rename to tests/prompting/prompt_control/test_prompt_control_expansion_probe.py diff --git a/tests/test_prompt_control_history_outputs.py b/tests/prompting/prompt_control/test_prompt_control_history_outputs.py similarity index 100% rename from tests/test_prompt_control_history_outputs.py rename to tests/prompting/prompt_control/test_prompt_control_history_outputs.py diff --git a/tests/test_prompt_control_prompt.py b/tests/prompting/prompt_control/test_prompt_control_prompt.py similarity index 100% rename from tests/test_prompt_control_prompt.py rename to tests/prompting/prompt_control/test_prompt_control_prompt.py diff --git a/tests/test_prompt_control_regional_hook_identities.py b/tests/prompting/prompt_control/test_prompt_control_regional_hook_identities.py similarity index 100% rename from tests/test_prompt_control_regional_hook_identities.py rename to tests/prompting/prompt_control/test_prompt_control_regional_hook_identities.py diff --git a/tests/test_prompt_control_results.py b/tests/prompting/prompt_control/test_prompt_control_results.py similarity index 89% rename from tests/test_prompt_control_results.py rename to tests/prompting/prompt_control/test_prompt_control_results.py index ca7f6f1..fa7a0d6 100644 --- a/tests/test_prompt_control_results.py +++ b/tests/prompting/prompt_control/test_prompt_control_results.py @@ -7,7 +7,6 @@ from pathlib import Path import pytest -from test_prompt_control_evidence_validation import _outputs from tools.prompt_control_characterization.cases import cases from tools.prompt_control_characterization.results import PromptControlResultRecorder @@ -15,6 +14,8 @@ from tools.prompt_control_characterization.source_identity import ( PromptControlSourceIdentity, ) +from .support.evidence import prompt_control_outputs + def test_recorder_finalizes_only_complete_successful_matrix(tmp_path: Path) -> None: """Persist one custom matrix in authoritative order and finalize it.""" @@ -28,7 +29,11 @@ def test_recorder_finalizes_only_complete_successful_matrix(tmp_path: Path) -> N ) with pytest.raises(ValueError, match="incomplete"): recorder.finalize() - recorder.record_success(case, _outputs(case), {"1": {"class_type": "Test"}}) + recorder.record_success( + case, + prompt_control_outputs(case), + {"1": {"class_type": "Test"}}, + ) result = recorder.finalize() assert result.is_file() assert case.case_id in recorder.completed_case_ids diff --git a/tests/test_prompt_control_runtime_probe.py b/tests/prompting/prompt_control/test_prompt_control_runtime_probe.py similarity index 100% rename from tests/test_prompt_control_runtime_probe.py rename to tests/prompting/prompt_control/test_prompt_control_runtime_probe.py diff --git a/tests/test_prompt_control_schedule_encode_graph.py b/tests/prompting/prompt_control/test_prompt_control_schedule_encode_graph.py similarity index 99% rename from tests/test_prompt_control_schedule_encode_graph.py rename to tests/prompting/prompt_control/test_prompt_control_schedule_encode_graph.py index 3b04b5f..4f51f37 100644 --- a/tests/test_prompt_control_schedule_encode_graph.py +++ b/tests/prompting/prompt_control/test_prompt_control_schedule_encode_graph.py @@ -13,7 +13,7 @@ from typing import Any, cast import pytest -from simple_syrup.runtime.prompt_control_schedule_encode_graph import ( +from simple_syrup.services.prompt_control_schedule_encode_graph import ( PROMPT_CONTROL_MISSING_MESSAGE, PromptControlScheduleEncodeGraphBuilder, ) diff --git a/tests/test_prompt_control_segment_planning_service.py b/tests/prompting/prompt_control/test_prompt_control_segment_planning_service.py similarity index 100% rename from tests/test_prompt_control_segment_planning_service.py rename to tests/prompting/prompt_control/test_prompt_control_segment_planning_service.py diff --git a/tests/test_prompt_control_snapshot_probe.py b/tests/prompting/prompt_control/test_prompt_control_snapshot_probe.py similarity index 100% rename from tests/test_prompt_control_snapshot_probe.py rename to tests/prompting/prompt_control/test_prompt_control_snapshot_probe.py diff --git a/tests/test_prompt_control_source_identity.py b/tests/prompting/prompt_control/test_prompt_control_source_identity.py similarity index 100% rename from tests/test_prompt_control_source_identity.py rename to tests/prompting/prompt_control/test_prompt_control_source_identity.py diff --git a/tests/test_prompt_encode_style_nodes.py b/tests/prompting/prompt_control/test_prompt_encode_style_nodes.py similarity index 100% rename from tests/test_prompt_encode_style_nodes.py rename to tests/prompting/prompt_control/test_prompt_encode_style_nodes.py diff --git a/tests/test_prompt_segment_alignment.py b/tests/prompting/prompt_control/test_prompt_segment_alignment.py similarity index 100% rename from tests/test_prompt_segment_alignment.py rename to tests/prompting/prompt_control/test_prompt_segment_alignment.py diff --git a/tests/test_schedule_and_encode_prompts_with_prompt_control_node.py b/tests/prompting/prompt_control/test_schedule_and_encode_prompts_with_prompt_control_node.py similarity index 100% rename from tests/test_schedule_and_encode_prompts_with_prompt_control_node.py rename to tests/prompting/prompt_control/test_schedule_and_encode_prompts_with_prompt_control_node.py diff --git a/tests/regional_generation/__init__.py b/tests/regional_generation/__init__.py new file mode 100644 index 0000000..bbdab81 --- /dev/null +++ b/tests/regional_generation/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation test behavior.""" diff --git a/tests/regional_generation/anima/__init__.py b/tests/regional_generation/anima/__init__.py new file mode 100644 index 0000000..390dec2 --- /dev/null +++ b/tests/regional_generation/anima/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation anima test behavior.""" diff --git a/tests/regional_generation/anima/support/__init__.py b/tests/regional_generation/anima/support/__init__.py new file mode 100644 index 0000000..62e3c94 --- /dev/null +++ b/tests/regional_generation/anima/support/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation anima support test behavior.""" diff --git a/tests/anima_attention_coupling_call_harness.py b/tests/regional_generation/anima/support/anima_attention_coupling_call_harness.py similarity index 98% rename from tests/anima_attention_coupling_call_harness.py rename to tests/regional_generation/anima/support/anima_attention_coupling_call_harness.py index 4d9af5c..eef772d 100644 --- a/tests/anima_attention_coupling_call_harness.py +++ b/tests/regional_generation/anima/support/anima_attention_coupling_call_harness.py @@ -10,11 +10,6 @@ from dataclasses import dataclass from typing import Any, cast import torch -from attention_coupling_invariant_contract import ( - AttentionCouplingCallHarness, - AttentionCouplingCallObservation, - AttentionCouplingCallScenario, -) from torch import nn from simple_syrup.domain.regional_attention_batch import ( @@ -31,6 +26,12 @@ from simple_syrup.runtime.regional_lora.anima_attention_execution import ( ) from simple_syrup.runtime.regional_lora.anima_module_surface import AnimaModuleSurface +from ...attention_coupling.support.attention_coupling_invariant_contract import ( + AttentionCouplingCallHarness, + AttentionCouplingCallObservation, + AttentionCouplingCallScenario, +) + class _AnimaDiffusionModel(nn.Module): """Provide weak-referenceable Anima model identity for call invariants.""" diff --git a/tests/anima_attention_coupling_diagnostics_harness.py b/tests/regional_generation/anima/support/anima_attention_coupling_diagnostics_harness.py similarity index 79% rename from tests/anima_attention_coupling_diagnostics_harness.py rename to tests/regional_generation/anima/support/anima_attention_coupling_diagnostics_harness.py index 3e36d0e..1dcd4f4 100644 --- a/tests/anima_attention_coupling_diagnostics_harness.py +++ b/tests/regional_generation/anima/support/anima_attention_coupling_diagnostics_harness.py @@ -6,12 +6,15 @@ from __future__ import annotations -from attention_coupling_diagnostics_harness import build_diagnostics_observation -from attention_coupling_invariant_contract import ( +from comfy.ldm.anima.model import Anima + +from ...attention_coupling.support.attention_coupling_diagnostics_harness import ( + build_diagnostics_observation, +) +from ...attention_coupling.support.attention_coupling_invariant_contract import ( AttentionCouplingDiagnosticsHarness, AttentionCouplingDiagnosticsObservation, ) -from comfy.ldm.anima.model import Anima class AnimaAttentionCouplingDiagnosticsHarness(AttentionCouplingDiagnosticsHarness): diff --git a/tests/anima_attention_coupling_invariant_harness.py b/tests/regional_generation/anima/support/anima_attention_coupling_invariant_harness.py similarity index 99% rename from tests/anima_attention_coupling_invariant_harness.py rename to tests/regional_generation/anima/support/anima_attention_coupling_invariant_harness.py index 23f8d56..a588961 100644 --- a/tests/anima_attention_coupling_invariant_harness.py +++ b/tests/regional_generation/anima/support/anima_attention_coupling_invariant_harness.py @@ -7,12 +7,6 @@ from __future__ import annotations import torch -from attention_coupling_invariant_contract import ( - AttentionCouplingInvariantEntry, - AttentionCouplingInvariantHarness, - AttentionCouplingInvariantObservation, - AttentionCouplingInvariantScenario, -) from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -44,6 +38,13 @@ from simple_syrup.runtime.regional_lora.anima_cross_attention_context import ( AnimaCrossAttentionInvocationContext, ) +from ...attention_coupling.support.attention_coupling_invariant_contract import ( + AttentionCouplingInvariantEntry, + AttentionCouplingInvariantHarness, + AttentionCouplingInvariantObservation, + AttentionCouplingInvariantScenario, +) + class _DeterministicAnimaAttention(nn.Module): """Evaluate one packed branch call using recognizable context values.""" diff --git a/tests/anima_attention_coupling_lifecycle_harness.py b/tests/regional_generation/anima/support/anima_attention_coupling_lifecycle_harness.py similarity index 89% rename from tests/anima_attention_coupling_lifecycle_harness.py rename to tests/regional_generation/anima/support/anima_attention_coupling_lifecycle_harness.py index c6b4551..f558409 100644 --- a/tests/anima_attention_coupling_lifecycle_harness.py +++ b/tests/regional_generation/anima/support/anima_attention_coupling_lifecycle_harness.py @@ -10,12 +10,6 @@ from copy import deepcopy from typing import Any import torch -from anima_attention_coupling_lifecycle_values import meta_anima_model -from attention_coupling_invariant_contract import ( - AttentionCouplingLifecycleHarness, - AttentionCouplingLifecycleObservation, -) -from attention_coupling_invariant_values import scheduled_invariant_plan from torch import nn from simple_syrup.domain.regional_lora_plan import EMPTY_REGIONAL_LORA_PLAN @@ -24,6 +18,17 @@ from simple_syrup.runtime.regional_lora.anima_full_context_backend import ( ) from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation +from ...attention_coupling.support.attention_coupling_invariant_contract import ( + AttentionCouplingLifecycleHarness, + AttentionCouplingLifecycleObservation, +) +from ...attention_coupling.support.attention_coupling_invariant_values import ( + scheduled_invariant_plan, +) +from .anima_attention_coupling_lifecycle_values import ( + meta_anima_model, +) + class _AnimaModelRoot(nn.Module): """Expose Anima at Comfy's diffusion-model patch root.""" diff --git a/tests/anima_attention_coupling_lifecycle_values.py b/tests/regional_generation/anima/support/anima_attention_coupling_lifecycle_values.py similarity index 100% rename from tests/anima_attention_coupling_lifecycle_values.py rename to tests/regional_generation/anima/support/anima_attention_coupling_lifecycle_values.py diff --git a/tests/anima_branch_test_values.py b/tests/regional_generation/anima/support/anima_branch_test_values.py similarity index 100% rename from tests/anima_branch_test_values.py rename to tests/regional_generation/anima/support/anima_branch_test_values.py diff --git a/tests/anima_lora_characterization_fixtures.py b/tests/regional_generation/anima/support/anima_lora_characterization_fixtures.py similarity index 100% rename from tests/anima_lora_characterization_fixtures.py rename to tests/regional_generation/anima/support/anima_lora_characterization_fixtures.py diff --git a/tests/anima_module_surface_fixtures.py b/tests/regional_generation/anima/support/anima_module_surface_fixtures.py similarity index 100% rename from tests/anima_module_surface_fixtures.py rename to tests/regional_generation/anima/support/anima_module_surface_fixtures.py diff --git a/tests/anima_nondiffusion_rejection_values.py b/tests/regional_generation/anima/support/anima_nondiffusion_rejection_values.py similarity index 100% rename from tests/anima_nondiffusion_rejection_values.py rename to tests/regional_generation/anima/support/anima_nondiffusion_rejection_values.py diff --git a/tests/test_anima_activation_context.py b/tests/regional_generation/anima/test_anima_activation_context.py similarity index 100% rename from tests/test_anima_activation_context.py rename to tests/regional_generation/anima/test_anima_activation_context.py diff --git a/tests/test_anima_adaln_pruning.py b/tests/regional_generation/anima/test_anima_adaln_pruning.py similarity index 98% rename from tests/test_anima_adaln_pruning.py rename to tests/regional_generation/anima/test_anima_adaln_pruning.py index 5782480..250cbd8 100644 --- a/tests/test_anima_adaln_pruning.py +++ b/tests/regional_generation/anima/test_anima_adaln_pruning.py @@ -7,7 +7,6 @@ from __future__ import annotations import torch -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -39,6 +38,10 @@ from simple_syrup.runtime.regional_lora.anima_query_activity import ( ) from simple_syrup.runtime.regional_lora.anima_query_masks import AnimaQueryMaskBatch +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _BranchAwareProjection(nn.Module): """Return recognizable values for every compact branch invocation.""" diff --git a/tests/test_anima_attention_context_wrapper.py b/tests/regional_generation/anima/test_anima_attention_context_wrapper.py similarity index 100% rename from tests/test_anima_attention_context_wrapper.py rename to tests/regional_generation/anima/test_anima_attention_context_wrapper.py diff --git a/tests/test_anima_attention_coupling_conditioning.py b/tests/regional_generation/anima/test_anima_attention_coupling_conditioning.py similarity index 100% rename from tests/test_anima_attention_coupling_conditioning.py rename to tests/regional_generation/anima/test_anima_attention_coupling_conditioning.py diff --git a/tests/test_anima_attention_coupling_integration.py b/tests/regional_generation/anima/test_anima_attention_coupling_integration.py similarity index 100% rename from tests/test_anima_attention_coupling_integration.py rename to tests/regional_generation/anima/test_anima_attention_coupling_integration.py diff --git a/tests/test_anima_attention_coupling_model_family.py b/tests/regional_generation/anima/test_anima_attention_coupling_model_family.py similarity index 100% rename from tests/test_anima_attention_coupling_model_family.py rename to tests/regional_generation/anima/test_anima_attention_coupling_model_family.py diff --git a/tests/test_anima_attention_device_cache_lifecycle.py b/tests/regional_generation/anima/test_anima_attention_device_cache_lifecycle.py similarity index 97% rename from tests/test_anima_attention_device_cache_lifecycle.py rename to tests/regional_generation/anima/test_anima_attention_device_cache_lifecycle.py index 49db1f1..b91aa21 100644 --- a/tests/test_anima_attention_device_cache_lifecycle.py +++ b/tests/regional_generation/anima/test_anima_attention_device_cache_lifecycle.py @@ -10,7 +10,6 @@ from typing import Any import torch from comfy.patcher_extension import CallbacksMP -from regional_attention_test_values import single_entry_regions from simple_syrup.domain.regional_attention import RegionalAttentionBranch from simple_syrup.domain.regional_attention_batch import ( @@ -34,6 +33,10 @@ from simple_syrup.runtime.regional_lora.anima_query_mask_context import ( AnimaQueryMaskContext, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + _DETACH_KEY = "simple_syrup.anima_attention_device_cache" diff --git a/tests/test_anima_attention_only_diagnostics.py b/tests/regional_generation/anima/test_anima_attention_only_diagnostics.py similarity index 100% rename from tests/test_anima_attention_only_diagnostics.py rename to tests/regional_generation/anima/test_anima_attention_only_diagnostics.py diff --git a/tests/test_anima_branch_batch.py b/tests/regional_generation/anima/test_anima_branch_batch.py similarity index 100% rename from tests/test_anima_branch_batch.py rename to tests/regional_generation/anima/test_anima_branch_batch.py diff --git a/tests/test_anima_composition_phase.py b/tests/regional_generation/anima/test_anima_composition_phase.py similarity index 100% rename from tests/test_anima_composition_phase.py rename to tests/regional_generation/anima/test_anima_composition_phase.py diff --git a/tests/test_anima_contextual_attention_coupling_integration.py b/tests/regional_generation/anima/test_anima_contextual_attention_coupling_integration.py similarity index 100% rename from tests/test_anima_contextual_attention_coupling_integration.py rename to tests/regional_generation/anima/test_anima_contextual_attention_coupling_integration.py diff --git a/tests/test_anima_cross_attention.py b/tests/regional_generation/anima/test_anima_cross_attention.py similarity index 99% rename from tests/test_anima_cross_attention.py rename to tests/regional_generation/anima/test_anima_cross_attention.py index 42eb691..2502d69 100644 --- a/tests/test_anima_cross_attention.py +++ b/tests/regional_generation/anima/test_anima_cross_attention.py @@ -14,7 +14,6 @@ import pytest import torch from comfy.ldm.anima.model import Anima from comfy.ldm.cosmos.predict2 import Attention, Block -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -65,6 +64,10 @@ from simple_syrup.runtime.regional_lora.anima_module_surface import ( ) from simple_syrup.runtime.regional_lora.anima_targets import ANIMA_BLOCK_COUNT +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _DeterministicCrossAttention(nn.Module): """Return context-owned constants while recording one complete branch call.""" diff --git a/tests/test_anima_cross_attention_weights.py b/tests/regional_generation/anima/test_anima_cross_attention_weights.py similarity index 100% rename from tests/test_anima_cross_attention_weights.py rename to tests/regional_generation/anima/test_anima_cross_attention_weights.py diff --git a/tests/test_anima_diagnostics_cache.py b/tests/regional_generation/anima/test_anima_diagnostics_cache.py similarity index 100% rename from tests/test_anima_diagnostics_cache.py rename to tests/regional_generation/anima/test_anima_diagnostics_cache.py diff --git a/tests/test_anima_diffusion_model_service.py b/tests/regional_generation/anima/test_anima_diffusion_model_service.py similarity index 100% rename from tests/test_anima_diffusion_model_service.py rename to tests/regional_generation/anima/test_anima_diffusion_model_service.py diff --git a/tests/test_anima_full_tile_lora_equivalence.py b/tests/regional_generation/anima/test_anima_full_tile_lora_equivalence.py similarity index 99% rename from tests/test_anima_full_tile_lora_equivalence.py rename to tests/regional_generation/anima/test_anima_full_tile_lora_equivalence.py index b1e973c..4a33690 100644 --- a/tests/test_anima_full_tile_lora_equivalence.py +++ b/tests/regional_generation/anima/test_anima_full_tile_lora_equivalence.py @@ -10,7 +10,6 @@ from uuid import uuid4 import pytest import torch -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -77,6 +76,10 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleSession, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _ZeroLinear(nn.Module): """Return an exact zero original output for isolated LoRA comparison.""" diff --git a/tests/test_anima_host_weight_interop.py b/tests/regional_generation/anima/test_anima_host_weight_interop.py similarity index 98% rename from tests/test_anima_host_weight_interop.py rename to tests/regional_generation/anima/test_anima_host_weight_interop.py index 87709d6..e56ae8c 100644 --- a/tests/test_anima_host_weight_interop.py +++ b/tests/regional_generation/anima/test_anima_host_weight_interop.py @@ -8,11 +8,6 @@ from __future__ import annotations import comfy.ops import torch -from regional_lora_test_values import ( - single_region_query_masks, - single_target_execution, - static_lora_schedule, -) from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch from simple_syrup.runtime.regional_lora.anima_composition import ( @@ -40,6 +35,12 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleResolution, ) +from ..regional.support.regional_lora_test_values import ( + single_region_query_masks, + single_target_execution, + static_lora_schedule, +) + def test_host_cast_function_composes_with_unchanged_regional_delta() -> None: """Apply Comfy's public global weight function before the regional delta.""" diff --git a/tests/test_anima_loader.py b/tests/regional_generation/anima/test_anima_loader.py similarity index 100% rename from tests/test_anima_loader.py rename to tests/regional_generation/anima/test_anima_loader.py diff --git a/tests/test_anima_lora_artifact_inventory.py b/tests/regional_generation/anima/test_anima_lora_artifact_inventory.py similarity index 100% rename from tests/test_anima_lora_artifact_inventory.py rename to tests/regional_generation/anima/test_anima_lora_artifact_inventory.py diff --git a/tests/test_anima_lora_block.py b/tests/regional_generation/anima/test_anima_lora_block.py similarity index 98% rename from tests/test_anima_lora_block.py rename to tests/regional_generation/anima/test_anima_lora_block.py index 1c2e657..c800aff 100644 --- a/tests/test_anima_lora_block.py +++ b/tests/regional_generation/anima/test_anima_lora_block.py @@ -11,8 +11,6 @@ from uuid import uuid4 import pytest import torch -from regional_attention_test_values import single_entry_regions -from regional_lora_test_values import static_lora_schedule from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -68,6 +66,13 @@ from simple_syrup.runtime.regional_lora.execution_cache import ( ) from simple_syrup.runtime.regional_lora.standard_adapter import StandardLoraTarget +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) +from ..regional.support.regional_lora_test_values import ( + static_lora_schedule, +) + class _CountingAttention(nn.Module): """Return one constant attention result and retain exact call count.""" diff --git a/tests/test_anima_lora_block_inactive_schedule.py b/tests/regional_generation/anima/test_anima_lora_block_inactive_schedule.py similarity index 98% rename from tests/test_anima_lora_block_inactive_schedule.py rename to tests/regional_generation/anima/test_anima_lora_block_inactive_schedule.py index 8eefaf8..076a797 100644 --- a/tests/test_anima_lora_block_inactive_schedule.py +++ b/tests/regional_generation/anima/test_anima_lora_block_inactive_schedule.py @@ -10,7 +10,6 @@ from typing import Any import pytest import torch -from regional_lora_test_values import single_target_execution from torch import nn from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch @@ -32,6 +31,10 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleResolution, ) +from ..regional.support.regional_lora_test_values import ( + single_target_execution, +) + class _PassAttention(nn.Module): """Count calls while preserving the supplied tensor.""" diff --git a/tests/test_anima_lora_characterization_cli.py b/tests/regional_generation/anima/test_anima_lora_characterization_cli.py similarity index 100% rename from tests/test_anima_lora_characterization_cli.py rename to tests/regional_generation/anima/test_anima_lora_characterization_cli.py diff --git a/tests/test_anima_lora_characterization_matrix.py b/tests/regional_generation/anima/test_anima_lora_characterization_matrix.py similarity index 100% rename from tests/test_anima_lora_characterization_matrix.py rename to tests/regional_generation/anima/test_anima_lora_characterization_matrix.py diff --git a/tests/test_anima_lora_characterization_workflow.py b/tests/regional_generation/anima/test_anima_lora_characterization_workflow.py similarity index 100% rename from tests/test_anima_lora_characterization_workflow.py rename to tests/regional_generation/anima/test_anima_lora_characterization_workflow.py diff --git a/tests/test_anima_lora_combined_multiplier_cache.py b/tests/regional_generation/anima/test_anima_lora_combined_multiplier_cache.py similarity index 98% rename from tests/test_anima_lora_combined_multiplier_cache.py rename to tests/regional_generation/anima/test_anima_lora_combined_multiplier_cache.py index 5f5d73d..43c7de5 100644 --- a/tests/test_anima_lora_combined_multiplier_cache.py +++ b/tests/regional_generation/anima/test_anima_lora_combined_multiplier_cache.py @@ -9,10 +9,6 @@ from __future__ import annotations from dataclasses import replace import torch -from regional_lora_test_values import ( - single_region_query_masks, - single_target_execution, -) from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch from simple_syrup.runtime.regional_lora.anima_composition import ( @@ -35,6 +31,11 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleResolution, ) +from ..regional.support.regional_lora_test_values import ( + single_region_query_masks, + single_target_execution, +) + def test_combined_multiplier_reuses_exact_scope_and_separates_strengths() -> None: """Cache the ordered sum while retaining schedule and mask-scope identity.""" diff --git a/tests/test_anima_lora_evidence_validation.py b/tests/regional_generation/anima/test_anima_lora_evidence_validation.py similarity index 95% rename from tests/test_anima_lora_evidence_validation.py rename to tests/regional_generation/anima/test_anima_lora_evidence_validation.py index 842d33e..a6a36d2 100644 --- a/tests/test_anima_lora_evidence_validation.py +++ b/tests/regional_generation/anima/test_anima_lora_evidence_validation.py @@ -9,7 +9,6 @@ from __future__ import annotations import copy import pytest -from anima_lora_characterization_fixtures import completed_outputs, inventory from tools.anima_lora_characterization.evidence_validation import ( validate_run_evidence, @@ -17,6 +16,11 @@ from tools.anima_lora_characterization.evidence_validation import ( from tools.anima_lora_characterization.history_outputs import LoraCompletedOutputs from tools.anima_lora_characterization.matrix import runs +from .support.anima_lora_characterization_fixtures import ( + completed_outputs, + inventory, +) + @pytest.mark.parametrize( "profile_id,capture_outputs", diff --git a/tests/test_anima_lora_history_outputs.py b/tests/regional_generation/anima/test_anima_lora_history_outputs.py similarity index 95% rename from tests/test_anima_lora_history_outputs.py rename to tests/regional_generation/anima/test_anima_lora_history_outputs.py index c0fc1b0..c9f9d89 100644 --- a/tests/test_anima_lora_history_outputs.py +++ b/tests/regional_generation/anima/test_anima_lora_history_outputs.py @@ -7,12 +7,15 @@ from __future__ import annotations import pytest -from anima_lora_characterization_fixtures import metrics from tools.anima_lora_characterization.history_outputs import parse_completed_outputs from tools.anima_lora_characterization.matrix import runs from tools.comfy_api import JsonObject +from .support.anima_lora_characterization_fixtures import ( + metrics, +) + def test_history_decoder_returns_one_probe_record_and_image() -> None: """Decode the exact output-node fields emitted by real Comfy history.""" diff --git a/tests/test_anima_lora_linear.py b/tests/regional_generation/anima/test_anima_lora_linear.py similarity index 99% rename from tests/test_anima_lora_linear.py rename to tests/regional_generation/anima/test_anima_lora_linear.py index 90a6e66..66ed016 100644 --- a/tests/test_anima_lora_linear.py +++ b/tests/regional_generation/anima/test_anima_lora_linear.py @@ -10,12 +10,6 @@ from dataclasses import replace import pytest import torch -from anima_branch_test_values import uniform_branch_invocation -from regional_lora_test_values import ( - single_region_query_masks, - single_target_execution, - static_lora_schedule, -) from torch import nn from simple_syrup.domain.regional_lora_plan import ( @@ -47,6 +41,15 @@ from simple_syrup.runtime.regional_lora.anima_targets import ( AnimaLoraTargetFamily, ) +from ..regional.support.regional_lora_test_values import ( + single_region_query_masks, + single_target_execution, + static_lora_schedule, +) +from .support.anima_branch_test_values import ( + uniform_branch_invocation, +) + _CROSS_FAMILIES = tuple( family for family in AnimaLoraTargetFamily if family.value.startswith("cross_attn") ) diff --git a/tests/test_anima_lora_results.py b/tests/regional_generation/anima/test_anima_lora_results.py similarity index 94% rename from tests/test_anima_lora_results.py rename to tests/regional_generation/anima/test_anima_lora_results.py index 7a656d3..2dd0f71 100644 --- a/tests/test_anima_lora_results.py +++ b/tests/regional_generation/anima/test_anima_lora_results.py @@ -10,12 +10,16 @@ import json from pathlib import Path import pytest -from anima_lora_characterization_fixtures import completed_outputs, inventory from tools.anima_lora_characterization.matrix import runs from tools.anima_lora_characterization.results import LoraResultRecorder from tools.comfy_api import JsonObject +from .support.anima_lora_characterization_fixtures import ( + completed_outputs, + inventory, +) + def test_recorder_persists_resumes_and_finalizes_exact_matrix( tmp_path: Path, diff --git a/tests/test_anima_lora_weight_categories.py b/tests/regional_generation/anima/test_anima_lora_weight_categories.py similarity index 100% rename from tests/test_anima_lora_weight_categories.py rename to tests/regional_generation/anima/test_anima_lora_weight_categories.py diff --git a/tests/test_anima_model_capability_detection.py b/tests/regional_generation/anima/test_anima_model_capability_detection.py similarity index 97% rename from tests/test_anima_model_capability_detection.py rename to tests/regional_generation/anima/test_anima_model_capability_detection.py index 4a10ca2..a5ed1b6 100644 --- a/tests/test_anima_model_capability_detection.py +++ b/tests/regional_generation/anima/test_anima_model_capability_detection.py @@ -14,7 +14,6 @@ import pytest import torch from comfy.ldm.anima.model import Anima as AnimaDiffusionModel from comfy.ldm.cosmos.predict2 import MiniTrainDIT -from regional_model_capability_test_values import empty_module, patcher from simple_syrup.domain.regional_model_capabilities import ( RegionalAttentionBackend, @@ -24,6 +23,11 @@ from simple_syrup.domain.regional_model_capabilities import ( ) from simple_syrup.runtime.anima_model_capability import AnimaModelCapabilityDetector +from ..regional.support.regional_model_capability_test_values import ( + empty_module, + patcher, +) + def test_exact_anima_surface_reports_specialized_capabilities() -> None: """Admit the installed Anima wrapper, diffusion module, and latent format.""" diff --git a/tests/test_anima_model_modifier_workflow.py b/tests/regional_generation/anima/test_anima_model_modifier_workflow.py similarity index 100% rename from tests/test_anima_model_modifier_workflow.py rename to tests/regional_generation/anima/test_anima_model_modifier_workflow.py diff --git a/tests/test_anima_model_patcher_surface.py b/tests/regional_generation/anima/test_anima_model_patcher_surface.py similarity index 98% rename from tests/test_anima_model_patcher_surface.py rename to tests/regional_generation/anima/test_anima_model_patcher_surface.py index b0677f0..04c33a4 100644 --- a/tests/test_anima_model_patcher_surface.py +++ b/tests/regional_generation/anima/test_anima_model_patcher_surface.py @@ -9,7 +9,6 @@ from __future__ import annotations from types import SimpleNamespace import pytest -from anima_module_surface_fixtures import installed_meta_anima from torch import nn from simple_syrup.runtime.model_patcher_mutations import ( @@ -26,6 +25,10 @@ from simple_syrup.runtime.regional_lora.anima_module_surface import ( ) from simple_syrup.runtime.regional_lora.anima_targets import ANIMA_BLOCK_COUNT +from .support.anima_module_surface_fixtures import ( + installed_meta_anima, +) + class _SourceVisiblePatcher: """Expose original objects while the shared live graph remains contaminated.""" diff --git a/tests/test_anima_model_reference.py b/tests/regional_generation/anima/test_anima_model_reference.py similarity index 100% rename from tests/test_anima_model_reference.py rename to tests/regional_generation/anima/test_anima_model_reference.py diff --git a/tests/test_anima_module_surface.py b/tests/regional_generation/anima/test_anima_module_surface.py similarity index 99% rename from tests/test_anima_module_surface.py rename to tests/regional_generation/anima/test_anima_module_surface.py index bd65ddb..10c4e53 100644 --- a/tests/test_anima_module_surface.py +++ b/tests/regional_generation/anima/test_anima_module_surface.py @@ -9,7 +9,6 @@ from __future__ import annotations from types import MethodType import pytest -from anima_module_surface_fixtures import installed_meta_anima from comfy.ldm.anima.model import Anima from comfy.ldm.cosmos.predict2 import Attention, Block, GPT2FeedForward from torch import nn @@ -28,6 +27,10 @@ from simple_syrup.runtime.regional_lora.anima_targets import ( expected_anima_lora_features, ) +from .support.anima_module_surface_fixtures import ( + installed_meta_anima, +) + @pytest.fixture(scope="module") def installed_anima() -> nn.Module: diff --git a/tests/test_anima_multi_lora_composition.py b/tests/regional_generation/anima/test_anima_multi_lora_composition.py similarity index 98% rename from tests/test_anima_multi_lora_composition.py rename to tests/regional_generation/anima/test_anima_multi_lora_composition.py index cce128b..44adab3 100644 --- a/tests/test_anima_multi_lora_composition.py +++ b/tests/regional_generation/anima/test_anima_multi_lora_composition.py @@ -10,9 +10,6 @@ from uuid import uuid4 import pytest import torch -from anima_branch_test_values import uniform_branch_invocation -from regional_attention_test_values import single_entry_regions -from regional_lora_test_values import static_lora_schedule from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -63,6 +60,16 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleResolution, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) +from ..regional.support.regional_lora_test_values import ( + static_lora_schedule, +) +from .support.anima_branch_test_values import ( + uniform_branch_invocation, +) + class _CountingZeroLinear(nn.Module): """Return a zero original output while counting exact calls.""" diff --git a/tests/test_anima_multi_lora_fidelity.py b/tests/regional_generation/anima/test_anima_multi_lora_fidelity.py similarity index 98% rename from tests/test_anima_multi_lora_fidelity.py rename to tests/regional_generation/anima/test_anima_multi_lora_fidelity.py index fad4b07..3bbde50 100644 --- a/tests/test_anima_multi_lora_fidelity.py +++ b/tests/regional_generation/anima/test_anima_multi_lora_fidelity.py @@ -9,10 +9,7 @@ from __future__ import annotations from uuid import uuid4 import torch -from anima_branch_test_values import uniform_branch_invocation from comfy.hooks import HookKeyframe, HookKeyframeGroup, WeightHook -from regional_attention_test_values import single_entry_regions -from regional_lora_test_values import complete_anima_admission, static_lora_schedule from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -67,6 +64,17 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleSession, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) +from ..regional.support.regional_lora_test_values import ( + complete_anima_admission, + static_lora_schedule, +) +from .support.anima_branch_test_values import ( + uniform_branch_invocation, +) + _SELF_FAMILIES = frozenset( family for family in AnimaLoraTargetFamily if family.value.startswith("self_attn") ) diff --git a/tests/test_anima_plain_comparison.py b/tests/regional_generation/anima/test_anima_plain_comparison.py similarity index 100% rename from tests/test_anima_plain_comparison.py rename to tests/regional_generation/anima/test_anima_plain_comparison.py diff --git a/tests/test_anima_projection_batch.py b/tests/regional_generation/anima/test_anima_projection_batch.py similarity index 100% rename from tests/test_anima_projection_batch.py rename to tests/regional_generation/anima/test_anima_projection_batch.py diff --git a/tests/test_anima_quantization_workflow.py b/tests/regional_generation/anima/test_anima_quantization_workflow.py similarity index 100% rename from tests/test_anima_quantization_workflow.py rename to tests/regional_generation/anima/test_anima_quantization_workflow.py diff --git a/tests/test_anima_query_activity.py b/tests/regional_generation/anima/test_anima_query_activity.py similarity index 98% rename from tests/test_anima_query_activity.py rename to tests/regional_generation/anima/test_anima_query_activity.py index ebe0b12..fb2d51c 100644 --- a/tests/test_anima_query_activity.py +++ b/tests/regional_generation/anima/test_anima_query_activity.py @@ -7,7 +7,6 @@ from __future__ import annotations import torch -from regional_attention_test_values import single_entry_regions from simple_syrup.domain.regional_attention import RegionalAttentionBranch from simple_syrup.domain.regional_attention_batch import ( @@ -37,6 +36,10 @@ from simple_syrup.runtime.regional_lora.anima_query_mask_context import ( ) from simple_syrup.runtime.regional_lora.anima_query_masks import AnimaQueryMaskBatch +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _FixedQueryMaskContext(AnimaQueryMaskContext): """Return one exact mixed-activity mask batch and record resolutions.""" diff --git a/tests/test_anima_query_mask_context.py b/tests/regional_generation/anima/test_anima_query_mask_context.py similarity index 98% rename from tests/test_anima_query_mask_context.py rename to tests/regional_generation/anima/test_anima_query_mask_context.py index 8193df1..b42c1ef 100644 --- a/tests/test_anima_query_mask_context.py +++ b/tests/regional_generation/anima/test_anima_query_mask_context.py @@ -7,7 +7,6 @@ from __future__ import annotations import torch -from regional_attention_test_values import single_entry_regions from simple_syrup.domain.regional_attention import RegionalAttentionBranch from simple_syrup.domain.regional_attention_batch import ( @@ -33,6 +32,10 @@ from simple_syrup.runtime.regional_lora.anima_query_masks import ( AnimaQueryMaskProjector, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _RecordingQueryMaskProjector(AnimaQueryMaskProjector): """Count exact canonical projections while retaining production behavior.""" diff --git a/tests/test_anima_query_masks.py b/tests/regional_generation/anima/test_anima_query_masks.py similarity index 100% rename from tests/test_anima_query_masks.py rename to tests/regional_generation/anima/test_anima_query_masks.py diff --git a/tests/regional_generation/anima/test_anima_regional_composition_diagnostics.py b/tests/regional_generation/anima/test_anima_regional_composition_diagnostics.py new file mode 100644 index 0000000..8f9c99e --- /dev/null +++ b/tests/regional_generation/anima/test_anima_regional_composition_diagnostics.py @@ -0,0 +1,399 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Prove exact structured diagnostics for regional Anima execution.""" + +from __future__ import annotations + +from uuid import uuid4 + +import comfy.ops +import pytest +import torch +from comfy.ldm.anima.model import Anima + +from simple_syrup.domain.regional_attention import RegionalAttentionBranch +from simple_syrup.domain.regional_attention_batch import ( + BatchedRegionalAttentionContexts, + BatchedRegionalAttentionEntry, + BatchedRegionalAttentionRegion, + RegionalAttentionChunkBatch, +) +from simple_syrup.domain.regional_lora_plan import ( + RegionalLoraAdapterIdentity, + RegionalLoraAdapterPlan, + RegionalLoraBranch, + RegionalLoraScheduleBoundary, +) +from simple_syrup.domain.regional_mask_bank import RegionalMaskBank +from simple_syrup.domain.spatial_views import ( + SpatialBatchLayout, +) +from simple_syrup.runtime.regional_lora.anima_activation_context import ( + AnimaActivationGeometry, +) +from simple_syrup.runtime.regional_lora.anima_attention_execution import ( + AnimaRegionalAttentionExecution, +) +from simple_syrup.runtime.regional_lora.anima_composition import ( + AnimaRegionalLoraComposition, +) +from simple_syrup.runtime.regional_lora.anima_composition_phase import ( + AnimaCompositionPhase, + AnimaCompositionStage, +) +from simple_syrup.runtime.regional_lora.anima_diagnostics import ( + AnimaRegionalDiagnosticsBuilder, +) +from simple_syrup.runtime.regional_lora.anima_execution_scope import ( + AnimaRegionalLoraAdapterExecution, +) +from simple_syrup.runtime.regional_lora.anima_module_surface import ( + ANIMA_MODULE_SURFACE_DISCOVERY, + AnimaLoraTargetModule, + AnimaModuleSurface, +) +from simple_syrup.runtime.regional_lora.anima_targets import ( + ANIMA_BLOCK_COUNT, + AnimaLoraAdmission, + AnimaLoraTarget, +) +from simple_syrup.runtime.regional_lora.execution_cache import ( + ModelCloneLineage, + RegionalLoraExecutionCache, +) +from simple_syrup.runtime.regional_lora.standard_adapter import StandardLoraTarget +from simple_syrup.runtime.regional_lora_schedule_resolution import ( + RegionalLoraScheduleResolution, +) + +from ..regional.support.regional_lora_test_values import ( + static_lora_schedule, +) + + +class _RecordingExecutor: + """Retain a verified class owner and forwarded call arguments.""" + + def __init__(self, class_obj: object, result: object) -> None: + """Store the exact owner and downstream result.""" + + self.class_obj = class_obj + self.result = result + self.calls: list[tuple[tuple[object, ...], dict[str, object]]] = [] + + def __call__(self, *args: object, **kwargs: object) -> object: + """Record one unchanged call and return the configured result.""" + + self.calls.append((args, kwargs)) + return self.result + + +@pytest.fixture(scope="module") +def anima_surface() -> AnimaModuleSurface: + """Discover a real installed Anima graph with allocation-free meta weights.""" + + model = Anima( + max_img_h=2, + max_img_w=2, + max_frames=1, + in_channels=16, + out_channels=16, + patch_spatial=2, + patch_temporal=1, + model_channels=2048, + num_blocks=ANIMA_BLOCK_COUNT, + num_heads=16, + mlp_ratio=4.0, + crossattn_emb_channels=1024, + pos_emb_cls="rope3d", + pos_emb_learnable=False, + pos_emb_interpolation="crop", + use_adaln_lora=True, + adaln_lora_dim=256, + extra_per_block_abs_pos_emb=False, + device=torch.device("meta"), + dtype=torch.float16, + operations=comfy.ops.disable_weight_init, + ) + return ANIMA_MODULE_SURFACE_DISCOVERY.discover(model) + + +def test_zero_pruned_composition_reports_no_low_rank_work( + anima_surface: AnimaModuleSurface, +) -> None: + """Distinguish regional attention work from a fully pruned adapter stack.""" + + composition = _composition( + anima_surface, + torch.zeros((1, 2, 2)), + include_negative=False, + all_zero_strength=True, + ) + + snapshot = AnimaRegionalDiagnosticsBuilder(anima_surface, composition).build( + _geometry(batch=1, height=2, width=2), + schedule_resolution=static_lora_schedule(composition)[1], + ) + work = snapshot.work + + assert snapshot.coverage_class == "all_base" + assert snapshot.active_region_indices == () + assert len(snapshot.adapter_uses) == 1 + assert snapshot.adapter_uses[0].active is False + assert snapshot.adapter_uses[0].pruning_reason == "zero_strength" + assert work.cross_attention_branch_multiplier == 2.0 + assert work.low_rank_adapter_multiplier == 0.0 + assert work.active_adapter_uses == 0 + assert work.active_target_count == 0 + assert work.target_use_count == 0 + assert work.compatible_projection_batch_count == 0 + assert work.denoiser_call_multiplier == 1.0 + + +def test_schedule_zero_reports_dynamic_pruning_without_hiding_static_support( + anima_surface: AnimaModuleSurface, +) -> None: + """Report current inactive adapter work separately from static pruning.""" + + composition = _composition( + anima_surface, + torch.ones((1, 2, 2)), + include_negative=False, + ) + count = len(composition.executions) + resolution = RegionalLoraScheduleResolution( + (0.0,) * count, + (0.0,) * count, + ) + + snapshot = AnimaRegionalDiagnosticsBuilder(anima_surface, composition).build( + _geometry(batch=1, height=2, width=2), + schedule_resolution=resolution, + ) + + assert all(not adapter.active for adapter in snapshot.adapter_uses) + assert all( + adapter.pruning_reason == "schedule_zero" for adapter in snapshot.adapter_uses + ) + assert snapshot.work.active_adapter_uses == 0 + assert snapshot.work.active_target_count == 0 + assert snapshot.work.target_use_count == 0 + + +def test_diagnostics_count_every_active_regional_conditioning_entry( + anima_surface: AnimaModuleSurface, +) -> None: + """Report entry-branch work after within-region schedule expansion.""" + + composition = _composition( + anima_surface, + torch.ones((1, 1, 1)), + include_negative=False, + regional_entry_count=2, + ) + + snapshot = AnimaRegionalDiagnosticsBuilder( + anima_surface, + composition, + ).build( + _geometry(batch=1, height=1, width=1), + schedule_resolution=static_lora_schedule(composition)[1], + ) + + assert snapshot.work.cross_attention_branch_multiplier == 3.0 + + +def test_repeated_chunk_layout_preserves_actual_order_and_latent_slices( + anima_surface: AnimaModuleSurface, +) -> None: + """Report repeated reversed CFG chunks and latent batches without assumptions.""" + + sequence = ( + RegionalAttentionBranch.NEGATIVE, + RegionalAttentionBranch.POSITIVE, + RegionalAttentionBranch.NEGATIVE, + RegionalAttentionBranch.POSITIVE, + ) + composition = _composition( + anima_surface, + torch.ones((1, 2, 2)), + include_negative=True, + branch_sequence=sequence, + latent_batch_size=2, + ) + + snapshot = AnimaRegionalDiagnosticsBuilder(anima_surface, composition).build( + _geometry(batch=8, height=2, width=2), + schedule_resolution=static_lora_schedule(composition)[1], + ) + + assert snapshot.active_branches == ("negative", "positive") + assert snapshot.positive_chunk_count == 2 + assert snapshot.negative_chunk_count == 2 + assert snapshot.latent_batch_size == 2 + assert [chunk.to_log_fields() for chunk in snapshot.chunks] == [ + { + "chunk_index": index, + "branch": branch.value, + "batch_start": index * 2, + "batch_stop": (index + 1) * 2, + } + for index, branch in enumerate(sequence) + ] + + +def _specialization_phase() -> AnimaCompositionPhase: + """Return stable coordinated phase values for wrapper diagnostics.""" + + return AnimaCompositionPhase( + AnimaCompositionStage.SPECIALIZATION, + 0.4, + True, + 0.9, + ) + + +def _composition( + surface: AnimaModuleSurface, + masks: torch.Tensor, + *, + include_negative: bool, + all_zero_strength: bool = False, + branch_sequence: tuple[RegionalAttentionBranch, ...] | None = None, + latent_batch_size: int = 1, + regional_entry_count: int = 1, +) -> AnimaRegionalLoraComposition: + """Build repeated and distinct adapter uses over one installed target path.""" + + region_count = int(masks.shape[0]) + sequence = branch_sequence or ( + ( + RegionalAttentionBranch.POSITIVE, + RegionalAttentionBranch.NEGATIVE, + ) + if include_negative + else (RegionalAttentionBranch.POSITIVE,) + ) + batch = len(sequence) * latent_batch_size + context = torch.zeros((batch, 1, 1)) + chunks = tuple( + RegionalAttentionChunkBatch( + index, + branch, + index * latent_batch_size, + (index + 1) * latent_batch_size, + ) + for index, branch in enumerate(sequence) + ) + attention = AnimaRegionalAttentionExecution( + BatchedRegionalAttentionContexts( + latent_batch_size, + chunks, + context, + tuple( + BatchedRegionalAttentionRegion( + region_index, + tuple( + BatchedRegionalAttentionEntry( + entry_index, + context.clone(), + (1.0,) * batch, + ) + for entry_index in range(regional_entry_count) + ), + ) + for region_index in range(region_count) + ), + ), + RegionalMaskBank( + masks.clone(), + masks.clone(), + int(masks.shape[-1]), + int(masks.shape[-2]), + ), + (1.0,) * region_count, + ) + descriptor = surface.lora_targets[0] + first_target = _target(descriptor, scalar=1.0) + second_target = _target(descriptor, scalar=-0.5) + model = ModelCloneLineage(uuid4(), uuid4()) + cache = RegionalLoraExecutionCache() + uses: tuple[tuple[AnimaLoraTarget, str, int, float], ...] = ( + (first_target, "private-adapter-path/one.safetensors", 0, 1.0), + ( + first_target, + "private-adapter-path/one.safetensors", + min(1, region_count - 1), + 0.5, + ), + ( + second_target, + "private-adapter-path/two.safetensors", + min(1, region_count - 1), + -0.25, + ), + ) + if all_zero_strength: + uses = ((first_target, "zero", 0, 0.0),) + executions = tuple( + AnimaRegionalLoraAdapterExecution( + RegionalLoraAdapterPlan( + RegionalLoraAdapterIdentity(identity), + composition_index=index, + region_index=region, + branch=( + RegionalLoraBranch.NEGATIVE + if include_negative and index == 2 + else RegionalLoraBranch.POSITIVE + ), + model_strength=strength, + schedule=(RegionalLoraScheduleBoundary(0.0, 100.0, 1.0, 0),), + ), + AnimaLoraAdmission((target,)), + attention, + model, + cache, + ) + for index, (target, identity, region, strength) in enumerate(uses) + ) + return AnimaRegionalLoraComposition(executions) + + +def _target(descriptor: AnimaLoraTargetModule, *, scalar: float) -> AnimaLoraTarget: + """Build one allocation-light rank-one target from a discovered descriptor.""" + + tensor = torch.tensor(scalar) + adapter = StandardLoraTarget( + descriptor.target_name, + tensor.expand(1, descriptor.input_features), + tensor.expand(descriptor.output_features, 1), + 1, + descriptor.input_features, + descriptor.output_features, + ) + return AnimaLoraTarget(descriptor.block_index, descriptor.family, adapter) + + +def _geometry( + *, + batch: int, + height: int, + width: int, + layout: SpatialBatchLayout | None = None, +) -> AnimaActivationGeometry: + """Build patch-size-one geometry for diagnostics-only tests.""" + + return AnimaActivationGeometry( + batch, + 1, + height, + width, + 1, + 1, + 1, + height, + width, + layout, + ) diff --git a/tests/test_anima_regional_diagnostics.py b/tests/regional_generation/anima/test_anima_regional_diagnostics.py similarity index 62% rename from tests/test_anima_regional_diagnostics.py rename to tests/regional_generation/anima/test_anima_regional_diagnostics.py index f903942..5807c9a 100644 --- a/tests/test_anima_regional_diagnostics.py +++ b/tests/regional_generation/anima/test_anima_regional_diagnostics.py @@ -7,15 +7,12 @@ from __future__ import annotations import json -import logging from uuid import uuid4 import comfy.ops import pytest import torch from comfy.ldm.anima.model import Anima -from regional_lora_test_values import static_lora_schedule -from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch from simple_syrup.domain.regional_attention_batch import ( @@ -37,7 +34,6 @@ from simple_syrup.domain.spatial_views import ( SpatialViewKind, ) from simple_syrup.runtime.regional_lora.anima_activation_context import ( - AnimaActivationContext, AnimaActivationGeometry, ) from simple_syrup.runtime.regional_lora.anima_attention_execution import ( @@ -50,15 +46,9 @@ from simple_syrup.runtime.regional_lora.anima_composition_phase import ( AnimaCompositionPhase, AnimaCompositionStage, ) -from simple_syrup.runtime.regional_lora.anima_composition_phase_context import ( - AnimaCompositionPhaseContext, -) from simple_syrup.runtime.regional_lora.anima_diagnostics import ( AnimaRegionalDiagnosticsBuilder, ) -from simple_syrup.runtime.regional_lora.anima_diagnostics_wrapper import ( - AnimaRegionalDiagnosticsDiffusionWrapper, -) from simple_syrup.runtime.regional_lora.anima_execution_scope import ( AnimaRegionalLoraAdapterExecution, ) @@ -82,6 +72,10 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleResolution, ) +from ..regional.support.regional_lora_test_values import ( + static_lora_schedule, +) + class _RecordingExecutor: """Retain a verified class owner and forwarded call arguments.""" @@ -339,266 +333,6 @@ def test_tile_layout_and_cfg_one_snapshot_use_resolved_execution_layout( ] -def test_zero_pruned_composition_reports_no_low_rank_work( - anima_surface: AnimaModuleSurface, -) -> None: - """Distinguish regional attention work from a fully pruned adapter stack.""" - - composition = _composition( - anima_surface, - torch.zeros((1, 2, 2)), - include_negative=False, - all_zero_strength=True, - ) - - snapshot = AnimaRegionalDiagnosticsBuilder(anima_surface, composition).build( - _geometry(batch=1, height=2, width=2), - schedule_resolution=static_lora_schedule(composition)[1], - ) - work = snapshot.work - - assert snapshot.coverage_class == "all_base" - assert snapshot.active_region_indices == () - assert len(snapshot.adapter_uses) == 1 - assert snapshot.adapter_uses[0].active is False - assert snapshot.adapter_uses[0].pruning_reason == "zero_strength" - assert work.cross_attention_branch_multiplier == 2.0 - assert work.low_rank_adapter_multiplier == 0.0 - assert work.active_adapter_uses == 0 - assert work.active_target_count == 0 - assert work.target_use_count == 0 - assert work.compatible_projection_batch_count == 0 - assert work.denoiser_call_multiplier == 1.0 - - -def test_schedule_zero_reports_dynamic_pruning_without_hiding_static_support( - anima_surface: AnimaModuleSurface, -) -> None: - """Report current inactive adapter work separately from static pruning.""" - - composition = _composition( - anima_surface, - torch.ones((1, 2, 2)), - include_negative=False, - ) - count = len(composition.executions) - resolution = RegionalLoraScheduleResolution( - (0.0,) * count, - (0.0,) * count, - ) - - snapshot = AnimaRegionalDiagnosticsBuilder(anima_surface, composition).build( - _geometry(batch=1, height=2, width=2), - schedule_resolution=resolution, - ) - - assert all(not adapter.active for adapter in snapshot.adapter_uses) - assert all( - adapter.pruning_reason == "schedule_zero" for adapter in snapshot.adapter_uses - ) - assert snapshot.work.active_adapter_uses == 0 - assert snapshot.work.active_target_count == 0 - assert snapshot.work.target_use_count == 0 - - -def test_diagnostics_count_every_active_regional_conditioning_entry( - anima_surface: AnimaModuleSurface, -) -> None: - """Report entry-branch work after within-region schedule expansion.""" - - composition = _composition( - anima_surface, - torch.ones((1, 1, 1)), - include_negative=False, - regional_entry_count=2, - ) - - snapshot = AnimaRegionalDiagnosticsBuilder( - anima_surface, - composition, - ).build( - _geometry(batch=1, height=1, width=1), - schedule_resolution=static_lora_schedule(composition)[1], - ) - - assert snapshot.work.cross_attention_branch_multiplier == 3.0 - - -def test_repeated_chunk_layout_preserves_actual_order_and_latent_slices( - anima_surface: AnimaModuleSurface, -) -> None: - """Report repeated reversed CFG chunks and latent batches without assumptions.""" - - sequence = ( - RegionalAttentionBranch.NEGATIVE, - RegionalAttentionBranch.POSITIVE, - RegionalAttentionBranch.NEGATIVE, - RegionalAttentionBranch.POSITIVE, - ) - composition = _composition( - anima_surface, - torch.ones((1, 2, 2)), - include_negative=True, - branch_sequence=sequence, - latent_batch_size=2, - ) - - snapshot = AnimaRegionalDiagnosticsBuilder(anima_surface, composition).build( - _geometry(batch=8, height=2, width=2), - schedule_resolution=static_lora_schedule(composition)[1], - ) - - assert snapshot.active_branches == ("negative", "positive") - assert snapshot.positive_chunk_count == 2 - assert snapshot.negative_chunk_count == 2 - assert snapshot.latent_batch_size == 2 - assert [chunk.to_log_fields() for chunk in snapshot.chunks] == [ - { - "chunk_index": index, - "branch": branch.value, - "batch_start": index * 2, - "batch_stop": (index + 1) * 2, - } - for index, branch in enumerate(sequence) - ] - - -def test_wrapper_emits_one_safe_record_and_preserves_downstream_call( - anima_surface: AnimaModuleSurface, - caplog: pytest.LogCaptureFixture, -) -> None: - """Emit one JSON-safe record without prompt, path, tensor, or output leakage.""" - - masks = torch.ones((1, 2, 2)) - composition = _composition(anima_surface, masks, include_negative=False) - context = AnimaActivationContext() - phase_context = AnimaCompositionPhaseContext() - schedule, resolution = static_lora_schedule(composition) - logger = logging.getLogger("test.anima.regional.diagnostics") - wrapper = AnimaRegionalDiagnosticsDiffusionWrapper( - anima_surface, - context, - AnimaRegionalDiagnosticsBuilder(anima_surface, composition), - schedule, - phase_context, - logger=logger, - ) - executor = _RecordingExecutor(anima_surface.diffusion_model, "prediction") - model_input = torch.zeros((1, 16, 1, 2, 2)) - secret_prompt = "never-log-this-prompt" - - with caplog.at_level(logging.INFO, logger=logger.name): - with ( - schedule.activate(resolution), - context.activate(_geometry(batch=1, height=2, width=2)), - phase_context.activate(_specialization_phase()), - ): - result = wrapper( - executor, - model_input, - transformer_options={"prompt": secret_prompt}, - ) - - records = [record for record in caplog.records if record.name == logger.name] - assert result == "prediction" - assert executor.calls == [ - ((model_input,), {"transformer_options": {"prompt": secret_prompt}}) - ] - assert len(records) == 1 - assert records[0].__dict__["operation"] == "anima_attention_coupling.execute" - serialized = json.dumps(records[0].__dict__["regional_diagnostics"], sort_keys=True) - assert secret_prompt not in serialized - assert "private-adapter-path" not in serialized - assert "tensor(" not in serialized - assert "prediction" not in serialized - fields = records[0].__dict__["regional_diagnostics"] - assert fields["composition_phase"] == { - "stage": "specialization", - "denoising_progress": 0.4, - "restrict_self_attention": True, - "regional_lora_scale": 0.9, - } - - -def test_wrapper_validates_without_building_disabled_info_payload( - anima_surface: AnimaModuleSurface, - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Avoid unused snapshot work while retaining dynamic call validation.""" - - composition = _composition( - anima_surface, - torch.ones((1, 2, 2)), - include_negative=False, - ) - context = AnimaActivationContext() - phase_context = AnimaCompositionPhaseContext() - schedule, resolution = static_lora_schedule(composition) - builder = AnimaRegionalDiagnosticsBuilder(anima_surface, composition) - logger = logging.getLogger("test.anima.regional.diagnostics.disabled") - logger.setLevel(logging.WARNING) - wrapper = AnimaRegionalDiagnosticsDiffusionWrapper( - anima_surface, - context, - builder, - schedule, - phase_context, - logger=logger, - ) - - def unexpected_build(*args: object, **kwargs: object) -> object: - """Fail if a disabled INFO payload is constructed.""" - - del args, kwargs - raise AssertionError("disabled diagnostics payload was built") - - monkeypatch.setattr(builder, "build", unexpected_build) - executor = _RecordingExecutor(anima_surface.diffusion_model, "prediction") - model_input = torch.zeros((1, 16, 1, 2, 2)) - - with ( - schedule.activate(resolution), - context.activate(_geometry(batch=1, height=2, width=2)), - ): - result = wrapper(executor, model_input, transformer_options={}) - - assert result == "prediction" - assert len(executor.calls) == 1 - - -def test_wrapper_fails_before_logging_for_wrong_model_or_missing_geometry( - anima_surface: AnimaModuleSurface, - caplog: pytest.LogCaptureFixture, -) -> None: - """Reject invalid wrapper ownership and lifetime without ambiguous records.""" - - composition = _composition( - anima_surface, - torch.ones((1, 2, 2)), - include_negative=False, - ) - context = AnimaActivationContext() - phase_context = AnimaCompositionPhaseContext() - schedule, _resolution = static_lora_schedule(composition) - logger = logging.getLogger("test.anima.regional.diagnostics.failures") - wrapper = AnimaRegionalDiagnosticsDiffusionWrapper( - anima_surface, - context, - AnimaRegionalDiagnosticsBuilder(anima_surface, composition), - schedule, - phase_context, - logger=logger, - ) - - with caplog.at_level(logging.INFO, logger=logger.name): - with pytest.raises(ValueError, match="does not own"): - wrapper(_RecordingExecutor(nn.Identity(), None)) - with pytest.raises(RuntimeError, match="unavailable outside"): - wrapper(_RecordingExecutor(anima_surface.diffusion_model, None)) - - assert [record for record in caplog.records if record.name == logger.name] == [] - - def _specialization_phase() -> AnimaCompositionPhase: """Return stable coordinated phase values for wrapper diagnostics.""" diff --git a/tests/regional_generation/anima/test_anima_regional_diagnostics_wrapper.py b/tests/regional_generation/anima/test_anima_regional_diagnostics_wrapper.py new file mode 100644 index 0000000..bbd8211 --- /dev/null +++ b/tests/regional_generation/anima/test_anima_regional_diagnostics_wrapper.py @@ -0,0 +1,418 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Prove exact structured diagnostics for regional Anima execution.""" + +from __future__ import annotations + +import json +import logging +from uuid import uuid4 + +import comfy.ops +import pytest +import torch +from comfy.ldm.anima.model import Anima +from torch import nn + +from simple_syrup.domain.regional_attention import RegionalAttentionBranch +from simple_syrup.domain.regional_attention_batch import ( + BatchedRegionalAttentionContexts, + BatchedRegionalAttentionEntry, + BatchedRegionalAttentionRegion, + RegionalAttentionChunkBatch, +) +from simple_syrup.domain.regional_lora_plan import ( + RegionalLoraAdapterIdentity, + RegionalLoraAdapterPlan, + RegionalLoraBranch, + RegionalLoraScheduleBoundary, +) +from simple_syrup.domain.regional_mask_bank import RegionalMaskBank +from simple_syrup.domain.spatial_views import ( + SpatialBatchLayout, +) +from simple_syrup.runtime.regional_lora.anima_activation_context import ( + AnimaActivationContext, + AnimaActivationGeometry, +) +from simple_syrup.runtime.regional_lora.anima_attention_execution import ( + AnimaRegionalAttentionExecution, +) +from simple_syrup.runtime.regional_lora.anima_composition import ( + AnimaRegionalLoraComposition, +) +from simple_syrup.runtime.regional_lora.anima_composition_phase import ( + AnimaCompositionPhase, + AnimaCompositionStage, +) +from simple_syrup.runtime.regional_lora.anima_composition_phase_context import ( + AnimaCompositionPhaseContext, +) +from simple_syrup.runtime.regional_lora.anima_diagnostics import ( + AnimaRegionalDiagnosticsBuilder, +) +from simple_syrup.runtime.regional_lora.anima_diagnostics_wrapper import ( + AnimaRegionalDiagnosticsDiffusionWrapper, +) +from simple_syrup.runtime.regional_lora.anima_execution_scope import ( + AnimaRegionalLoraAdapterExecution, +) +from simple_syrup.runtime.regional_lora.anima_module_surface import ( + ANIMA_MODULE_SURFACE_DISCOVERY, + AnimaLoraTargetModule, + AnimaModuleSurface, +) +from simple_syrup.runtime.regional_lora.anima_targets import ( + ANIMA_BLOCK_COUNT, + AnimaLoraAdmission, + AnimaLoraTarget, +) +from simple_syrup.runtime.regional_lora.execution_cache import ( + ModelCloneLineage, + RegionalLoraExecutionCache, +) +from simple_syrup.runtime.regional_lora.standard_adapter import StandardLoraTarget + +from ..regional.support.regional_lora_test_values import ( + static_lora_schedule, +) + + +class _RecordingExecutor: + """Retain a verified class owner and forwarded call arguments.""" + + def __init__(self, class_obj: object, result: object) -> None: + """Store the exact owner and downstream result.""" + + self.class_obj = class_obj + self.result = result + self.calls: list[tuple[tuple[object, ...], dict[str, object]]] = [] + + def __call__(self, *args: object, **kwargs: object) -> object: + """Record one unchanged call and return the configured result.""" + + self.calls.append((args, kwargs)) + return self.result + + +@pytest.fixture(scope="module") +def anima_surface() -> AnimaModuleSurface: + """Discover a real installed Anima graph with allocation-free meta weights.""" + + model = Anima( + max_img_h=2, + max_img_w=2, + max_frames=1, + in_channels=16, + out_channels=16, + patch_spatial=2, + patch_temporal=1, + model_channels=2048, + num_blocks=ANIMA_BLOCK_COUNT, + num_heads=16, + mlp_ratio=4.0, + crossattn_emb_channels=1024, + pos_emb_cls="rope3d", + pos_emb_learnable=False, + pos_emb_interpolation="crop", + use_adaln_lora=True, + adaln_lora_dim=256, + extra_per_block_abs_pos_emb=False, + device=torch.device("meta"), + dtype=torch.float16, + operations=comfy.ops.disable_weight_init, + ) + return ANIMA_MODULE_SURFACE_DISCOVERY.discover(model) + + +def test_wrapper_emits_one_safe_record_and_preserves_downstream_call( + anima_surface: AnimaModuleSurface, + caplog: pytest.LogCaptureFixture, +) -> None: + """Emit one JSON-safe record without prompt, path, tensor, or output leakage.""" + + masks = torch.ones((1, 2, 2)) + composition = _composition(anima_surface, masks, include_negative=False) + context = AnimaActivationContext() + phase_context = AnimaCompositionPhaseContext() + schedule, resolution = static_lora_schedule(composition) + logger = logging.getLogger("test.anima.regional.diagnostics") + wrapper = AnimaRegionalDiagnosticsDiffusionWrapper( + anima_surface, + context, + AnimaRegionalDiagnosticsBuilder(anima_surface, composition), + schedule, + phase_context, + logger=logger, + ) + executor = _RecordingExecutor(anima_surface.diffusion_model, "prediction") + model_input = torch.zeros((1, 16, 1, 2, 2)) + secret_prompt = "never-log-this-prompt" + + with caplog.at_level(logging.INFO, logger=logger.name): + with ( + schedule.activate(resolution), + context.activate(_geometry(batch=1, height=2, width=2)), + phase_context.activate(_specialization_phase()), + ): + result = wrapper( + executor, + model_input, + transformer_options={"prompt": secret_prompt}, + ) + + records = [record for record in caplog.records if record.name == logger.name] + assert result == "prediction" + assert executor.calls == [ + ((model_input,), {"transformer_options": {"prompt": secret_prompt}}) + ] + assert len(records) == 1 + assert records[0].__dict__["operation"] == "anima_attention_coupling.execute" + serialized = json.dumps(records[0].__dict__["regional_diagnostics"], sort_keys=True) + assert secret_prompt not in serialized + assert "private-adapter-path" not in serialized + assert "tensor(" not in serialized + assert "prediction" not in serialized + fields = records[0].__dict__["regional_diagnostics"] + assert fields["composition_phase"] == { + "stage": "specialization", + "denoising_progress": 0.4, + "restrict_self_attention": True, + "regional_lora_scale": 0.9, + } + + +def test_wrapper_validates_without_building_disabled_info_payload( + anima_surface: AnimaModuleSurface, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Avoid unused snapshot work while retaining dynamic call validation.""" + + composition = _composition( + anima_surface, + torch.ones((1, 2, 2)), + include_negative=False, + ) + context = AnimaActivationContext() + phase_context = AnimaCompositionPhaseContext() + schedule, resolution = static_lora_schedule(composition) + builder = AnimaRegionalDiagnosticsBuilder(anima_surface, composition) + logger = logging.getLogger("test.anima.regional.diagnostics.disabled") + logger.setLevel(logging.WARNING) + wrapper = AnimaRegionalDiagnosticsDiffusionWrapper( + anima_surface, + context, + builder, + schedule, + phase_context, + logger=logger, + ) + + def unexpected_build(*args: object, **kwargs: object) -> object: + """Fail if a disabled INFO payload is constructed.""" + + del args, kwargs + raise AssertionError("disabled diagnostics payload was built") + + monkeypatch.setattr(builder, "build", unexpected_build) + executor = _RecordingExecutor(anima_surface.diffusion_model, "prediction") + model_input = torch.zeros((1, 16, 1, 2, 2)) + + with ( + schedule.activate(resolution), + context.activate(_geometry(batch=1, height=2, width=2)), + ): + result = wrapper(executor, model_input, transformer_options={}) + + assert result == "prediction" + assert len(executor.calls) == 1 + + +def test_wrapper_fails_before_logging_for_wrong_model_or_missing_geometry( + anima_surface: AnimaModuleSurface, + caplog: pytest.LogCaptureFixture, +) -> None: + """Reject invalid wrapper ownership and lifetime without ambiguous records.""" + + composition = _composition( + anima_surface, + torch.ones((1, 2, 2)), + include_negative=False, + ) + context = AnimaActivationContext() + phase_context = AnimaCompositionPhaseContext() + schedule, _resolution = static_lora_schedule(composition) + logger = logging.getLogger("test.anima.regional.diagnostics.failures") + wrapper = AnimaRegionalDiagnosticsDiffusionWrapper( + anima_surface, + context, + AnimaRegionalDiagnosticsBuilder(anima_surface, composition), + schedule, + phase_context, + logger=logger, + ) + + with caplog.at_level(logging.INFO, logger=logger.name): + with pytest.raises(ValueError, match="does not own"): + wrapper(_RecordingExecutor(nn.Identity(), None)) + with pytest.raises(RuntimeError, match="unavailable outside"): + wrapper(_RecordingExecutor(anima_surface.diffusion_model, None)) + + assert [record for record in caplog.records if record.name == logger.name] == [] + + +def _specialization_phase() -> AnimaCompositionPhase: + """Return stable coordinated phase values for wrapper diagnostics.""" + + return AnimaCompositionPhase( + AnimaCompositionStage.SPECIALIZATION, + 0.4, + True, + 0.9, + ) + + +def _composition( + surface: AnimaModuleSurface, + masks: torch.Tensor, + *, + include_negative: bool, + all_zero_strength: bool = False, + branch_sequence: tuple[RegionalAttentionBranch, ...] | None = None, + latent_batch_size: int = 1, + regional_entry_count: int = 1, +) -> AnimaRegionalLoraComposition: + """Build repeated and distinct adapter uses over one installed target path.""" + + region_count = int(masks.shape[0]) + sequence = branch_sequence or ( + ( + RegionalAttentionBranch.POSITIVE, + RegionalAttentionBranch.NEGATIVE, + ) + if include_negative + else (RegionalAttentionBranch.POSITIVE,) + ) + batch = len(sequence) * latent_batch_size + context = torch.zeros((batch, 1, 1)) + chunks = tuple( + RegionalAttentionChunkBatch( + index, + branch, + index * latent_batch_size, + (index + 1) * latent_batch_size, + ) + for index, branch in enumerate(sequence) + ) + attention = AnimaRegionalAttentionExecution( + BatchedRegionalAttentionContexts( + latent_batch_size, + chunks, + context, + tuple( + BatchedRegionalAttentionRegion( + region_index, + tuple( + BatchedRegionalAttentionEntry( + entry_index, + context.clone(), + (1.0,) * batch, + ) + for entry_index in range(regional_entry_count) + ), + ) + for region_index in range(region_count) + ), + ), + RegionalMaskBank( + masks.clone(), + masks.clone(), + int(masks.shape[-1]), + int(masks.shape[-2]), + ), + (1.0,) * region_count, + ) + descriptor = surface.lora_targets[0] + first_target = _target(descriptor, scalar=1.0) + second_target = _target(descriptor, scalar=-0.5) + model = ModelCloneLineage(uuid4(), uuid4()) + cache = RegionalLoraExecutionCache() + uses: tuple[tuple[AnimaLoraTarget, str, int, float], ...] = ( + (first_target, "private-adapter-path/one.safetensors", 0, 1.0), + ( + first_target, + "private-adapter-path/one.safetensors", + min(1, region_count - 1), + 0.5, + ), + ( + second_target, + "private-adapter-path/two.safetensors", + min(1, region_count - 1), + -0.25, + ), + ) + if all_zero_strength: + uses = ((first_target, "zero", 0, 0.0),) + executions = tuple( + AnimaRegionalLoraAdapterExecution( + RegionalLoraAdapterPlan( + RegionalLoraAdapterIdentity(identity), + composition_index=index, + region_index=region, + branch=( + RegionalLoraBranch.NEGATIVE + if include_negative and index == 2 + else RegionalLoraBranch.POSITIVE + ), + model_strength=strength, + schedule=(RegionalLoraScheduleBoundary(0.0, 100.0, 1.0, 0),), + ), + AnimaLoraAdmission((target,)), + attention, + model, + cache, + ) + for index, (target, identity, region, strength) in enumerate(uses) + ) + return AnimaRegionalLoraComposition(executions) + + +def _target(descriptor: AnimaLoraTargetModule, *, scalar: float) -> AnimaLoraTarget: + """Build one allocation-light rank-one target from a discovered descriptor.""" + + tensor = torch.tensor(scalar) + adapter = StandardLoraTarget( + descriptor.target_name, + tensor.expand(1, descriptor.input_features), + tensor.expand(descriptor.output_features, 1), + 1, + descriptor.input_features, + descriptor.output_features, + ) + return AnimaLoraTarget(descriptor.block_index, descriptor.family, adapter) + + +def _geometry( + *, + batch: int, + height: int, + width: int, + layout: SpatialBatchLayout | None = None, +) -> AnimaActivationGeometry: + """Build patch-size-one geometry for diagnostics-only tests.""" + + return AnimaActivationGeometry( + batch, + 1, + height, + width, + 1, + 1, + 1, + height, + width, + layout, + ) diff --git a/tests/test_anima_regional_lora_admission_history.py b/tests/regional_generation/anima/test_anima_regional_lora_admission_history.py similarity index 100% rename from tests/test_anima_regional_lora_admission_history.py rename to tests/regional_generation/anima/test_anima_regional_lora_admission_history.py diff --git a/tests/test_anima_regional_lora_admission_results.py b/tests/regional_generation/anima/test_anima_regional_lora_admission_results.py similarity index 100% rename from tests/test_anima_regional_lora_admission_results.py rename to tests/regional_generation/anima/test_anima_regional_lora_admission_results.py diff --git a/tests/test_anima_regional_lora_admission_workflow.py b/tests/regional_generation/anima/test_anima_regional_lora_admission_workflow.py similarity index 100% rename from tests/test_anima_regional_lora_admission_workflow.py rename to tests/regional_generation/anima/test_anima_regional_lora_admission_workflow.py diff --git a/tests/test_anima_regional_lora_attention_backend.py b/tests/regional_generation/anima/test_anima_regional_lora_attention_backend.py similarity index 100% rename from tests/test_anima_regional_lora_attention_backend.py rename to tests/regional_generation/anima/test_anima_regional_lora_attention_backend.py diff --git a/tests/test_anima_regional_lora_device_cache_lifecycle.py b/tests/regional_generation/anima/test_anima_regional_lora_device_cache_lifecycle.py similarity index 99% rename from tests/test_anima_regional_lora_device_cache_lifecycle.py rename to tests/regional_generation/anima/test_anima_regional_lora_device_cache_lifecycle.py index 753ff93..f47762a 100644 --- a/tests/test_anima_regional_lora_device_cache_lifecycle.py +++ b/tests/regional_generation/anima/test_anima_regional_lora_device_cache_lifecycle.py @@ -11,11 +11,6 @@ from typing import Any import torch from comfy.patcher_extension import CallbacksMP -from regional_lora_test_values import ( - single_region_query_masks, - single_target_execution, - static_lora_schedule, -) from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch from simple_syrup.runtime.regional_lora.active_support import ( @@ -50,6 +45,12 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleResolution, ) +from ..regional.support.regional_lora_test_values import ( + single_region_query_masks, + single_target_execution, + static_lora_schedule, +) + _DETACH_KEY = "simple_syrup.anima_regional_lora_device_cache" diff --git a/tests/test_anima_regional_lora_isolated_measurement.py b/tests/regional_generation/anima/test_anima_regional_lora_isolated_measurement.py similarity index 100% rename from tests/test_anima_regional_lora_isolated_measurement.py rename to tests/regional_generation/anima/test_anima_regional_lora_isolated_measurement.py diff --git a/tests/test_anima_regional_lora_isolated_runner.py b/tests/regional_generation/anima/test_anima_regional_lora_isolated_runner.py similarity index 100% rename from tests/test_anima_regional_lora_isolated_runner.py rename to tests/regional_generation/anima/test_anima_regional_lora_isolated_runner.py diff --git a/tests/test_anima_regional_lora_isolated_suite.py b/tests/regional_generation/anima/test_anima_regional_lora_isolated_suite.py similarity index 100% rename from tests/test_anima_regional_lora_isolated_suite.py rename to tests/regional_generation/anima/test_anima_regional_lora_isolated_suite.py diff --git a/tests/test_anima_regional_lora_operator_capture.py b/tests/regional_generation/anima/test_anima_regional_lora_operator_capture.py similarity index 100% rename from tests/test_anima_regional_lora_operator_capture.py rename to tests/regional_generation/anima/test_anima_regional_lora_operator_capture.py diff --git a/tests/test_anima_regional_lora_performance_artifacts.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_artifacts.py similarity index 100% rename from tests/test_anima_regional_lora_performance_artifacts.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_artifacts.py diff --git a/tests/test_anima_regional_lora_performance_manifest.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_manifest.py similarity index 100% rename from tests/test_anima_regional_lora_performance_manifest.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_manifest.py diff --git a/tests/test_anima_regional_lora_performance_measurement.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_measurement.py similarity index 100% rename from tests/test_anima_regional_lora_performance_measurement.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_measurement.py diff --git a/tests/test_anima_regional_lora_performance_profile.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_profile.py similarity index 98% rename from tests/test_anima_regional_lora_performance_profile.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_profile.py index c991369..05c4b1a 100644 --- a/tests/test_anima_regional_lora_performance_profile.py +++ b/tests/regional_generation/anima/test_anima_regional_lora_performance_profile.py @@ -10,7 +10,6 @@ from pathlib import Path from typing import Any, cast import pytest -from regional_lora_test_values import single_target_execution from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch from simple_syrup.runtime.regional_lora.anima_module_surface import AnimaModuleSurface @@ -22,6 +21,10 @@ from tools.anima_regional_lora_performance.runtime_profile import ( PerformanceRuntimeProfile, ) +from ..regional.support.regional_lora_test_values import ( + single_target_execution, +) + def test_p57_profile_builds_ordered_shared_cache_adapter_uses( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_anima_regional_lora_performance_profiling.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_profiling.py similarity index 100% rename from tests/test_anima_regional_lora_performance_profiling.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_profiling.py diff --git a/tests/test_anima_regional_lora_performance_results.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_results.py similarity index 100% rename from tests/test_anima_regional_lora_performance_results.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_results.py diff --git a/tests/test_anima_regional_lora_performance_runner.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_runner.py similarity index 100% rename from tests/test_anima_regional_lora_performance_runner.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_runner.py diff --git a/tests/test_anima_regional_lora_performance_suite.py b/tests/regional_generation/anima/test_anima_regional_lora_performance_suite.py similarity index 100% rename from tests/test_anima_regional_lora_performance_suite.py rename to tests/regional_generation/anima/test_anima_regional_lora_performance_suite.py diff --git a/tests/test_anima_regional_lora_plan_admission.py b/tests/regional_generation/anima/test_anima_regional_lora_plan_admission.py similarity index 100% rename from tests/test_anima_regional_lora_plan_admission.py rename to tests/regional_generation/anima/test_anima_regional_lora_plan_admission.py diff --git a/tests/test_anima_regional_lora_runtime_profile.py b/tests/regional_generation/anima/test_anima_regional_lora_runtime_profile.py similarity index 98% rename from tests/test_anima_regional_lora_runtime_profile.py rename to tests/regional_generation/anima/test_anima_regional_lora_runtime_profile.py index b136f8d..462e8c4 100644 --- a/tests/test_anima_regional_lora_runtime_profile.py +++ b/tests/regional_generation/anima/test_anima_regional_lora_runtime_profile.py @@ -12,7 +12,6 @@ from uuid import UUID, uuid4 import comfy.sampler_helpers import pytest -from regional_lora_test_values import single_target_execution from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch from simple_syrup.runtime.regional_lora.anima_composition import ( @@ -27,6 +26,10 @@ from tools.anima_regional_lora_performance.runtime_profile import ( PerformanceRegionalAdapterUse, ) +from ..regional.support.regional_lora_test_values import ( + single_target_execution, +) + @dataclass(frozen=True, slots=True) class _SourceModel: diff --git a/tests/test_anima_regional_lora_scaling_cli.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_cli.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_cli.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_cli.py diff --git a/tests/test_anima_regional_lora_scaling_fixture.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_fixture.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_fixture.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_fixture.py diff --git a/tests/test_anima_regional_lora_scaling_manifest.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_manifest.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_manifest.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_manifest.py diff --git a/tests/test_anima_regional_lora_scaling_measurement.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_measurement.py similarity index 97% rename from tests/test_anima_regional_lora_scaling_measurement.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_measurement.py index c1ea478..483d098 100644 --- a/tests/test_anima_regional_lora_scaling_measurement.py +++ b/tests/regional_generation/anima/test_anima_regional_lora_scaling_measurement.py @@ -10,7 +10,6 @@ from typing import Any, cast import pytest import torch -from regional_lora_test_values import single_target_execution from torch import nn from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch @@ -38,6 +37,10 @@ from tools.anima_regional_lora_performance.runtime_profile import ( PerformanceRuntimeProfile, ) +from ..regional.support.regional_lora_test_values import ( + single_target_execution, +) + def test_scaling_measurement_reads_production_work_and_cache( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_anima_regional_lora_scaling_profile.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_profile.py similarity index 98% rename from tests/test_anima_regional_lora_scaling_profile.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_profile.py index 006402c..7fa801f 100644 --- a/tests/test_anima_regional_lora_scaling_profile.py +++ b/tests/regional_generation/anima/test_anima_regional_lora_scaling_profile.py @@ -10,7 +10,6 @@ from pathlib import Path from typing import Any, cast import pytest -from regional_lora_test_values import single_target_execution from simple_syrup.domain.regional_lora_plan import RegionalLoraBranch from simple_syrup.runtime.regional_lora.anima_module_surface import AnimaModuleSurface @@ -27,6 +26,10 @@ from tools.anima_regional_lora_performance.runtime_profile import ( PerformanceRuntimeProfile, ) +from ..regional.support.regional_lora_test_values import ( + single_target_execution, +) + @pytest.mark.parametrize( ("profile_index", "regions", "strengths", "identities", "schedule"), diff --git a/tests/test_anima_regional_lora_scaling_profiling.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_profiling.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_profiling.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_profiling.py diff --git a/tests/test_anima_regional_lora_scaling_results.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_results.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_results.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_results.py diff --git a/tests/test_anima_regional_lora_scaling_runner.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_runner.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_runner.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_runner.py diff --git a/tests/test_anima_regional_lora_scaling_suite.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_suite.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_suite.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_suite.py diff --git a/tests/test_anima_regional_lora_scaling_work_contract.py b/tests/regional_generation/anima/test_anima_regional_lora_scaling_work_contract.py similarity index 100% rename from tests/test_anima_regional_lora_scaling_work_contract.py rename to tests/regional_generation/anima/test_anima_regional_lora_scaling_work_contract.py diff --git a/tests/test_anima_regional_lora_target_ownership.py b/tests/regional_generation/anima/test_anima_regional_lora_target_ownership.py similarity index 100% rename from tests/test_anima_regional_lora_target_ownership.py rename to tests/regional_generation/anima/test_anima_regional_lora_target_ownership.py diff --git a/tests/test_anima_regional_lora_targets.py b/tests/regional_generation/anima/test_anima_regional_lora_targets.py similarity index 100% rename from tests/test_anima_regional_lora_targets.py rename to tests/regional_generation/anima/test_anima_regional_lora_targets.py diff --git a/tests/test_anima_regional_lora_vram_cli.py b/tests/regional_generation/anima/test_anima_regional_lora_vram_cli.py similarity index 100% rename from tests/test_anima_regional_lora_vram_cli.py rename to tests/regional_generation/anima/test_anima_regional_lora_vram_cli.py diff --git a/tests/test_anima_regional_model_smoke.py b/tests/regional_generation/anima/test_anima_regional_model_smoke.py similarity index 98% rename from tests/test_anima_regional_model_smoke.py rename to tests/regional_generation/anima/test_anima_regional_model_smoke.py index dc6ee71..0d7746e 100644 --- a/tests/test_anima_regional_model_smoke.py +++ b/tests/regional_generation/anima/test_anima_regional_model_smoke.py @@ -14,7 +14,6 @@ import comfy.sampler_helpers import pytest import torch from comfy.ldm.anima.model import Anima -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -57,6 +56,10 @@ from simple_syrup.runtime.regional_lora.execution_cache import ( ) from simple_syrup.runtime.regional_lora.standard_adapter import StandardLoraTarget +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _AnimaModelRoot(nn.Module): """Expose installed Anima under Comfy's diffusion-model patch path.""" diff --git a/tests/test_anima_regional_nondiffusion_rejection_matrix.py b/tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_matrix.py similarity index 100% rename from tests/test_anima_regional_nondiffusion_rejection_matrix.py rename to tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_matrix.py diff --git a/tests/test_anima_regional_nondiffusion_rejection_results.py b/tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_results.py similarity index 97% rename from tests/test_anima_regional_nondiffusion_rejection_results.py rename to tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_results.py index ae920c0..3a30c51 100644 --- a/tests/test_anima_regional_nondiffusion_rejection_results.py +++ b/tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_results.py @@ -10,7 +10,6 @@ import json from pathlib import Path import pytest -from anima_nondiffusion_rejection_values import synthetic_rejection from tools.anima_regional_lora_admission_integration.workflow import ( RegionalLoraAdmissionWorkflowBuilder, @@ -20,6 +19,10 @@ from tools.anima_regional_nondiffusion_rejection_integration.results import ( AnimaNondiffusionRejectionResultRecorder, ) +from .support.anima_nondiffusion_rejection_values import ( + synthetic_rejection, +) + def test_recorder_publishes_four_labeled_rejections_and_no_image( tmp_path: Path, diff --git a/tests/test_anima_regional_nondiffusion_rejection_validation.py b/tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_validation.py similarity index 96% rename from tests/test_anima_regional_nondiffusion_rejection_validation.py rename to tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_validation.py index b8b77da..6024cd9 100644 --- a/tests/test_anima_regional_nondiffusion_rejection_validation.py +++ b/tests/regional_generation/anima/test_anima_regional_nondiffusion_rejection_validation.py @@ -7,7 +7,6 @@ from __future__ import annotations import pytest -from anima_nondiffusion_rejection_values import synthetic_rejection_history from tools.anima_regional_lora_admission_integration.history import ( RegionalLoraAdmissionError, @@ -21,6 +20,10 @@ from tools.anima_regional_nondiffusion_rejection_integration.validation import ( validate_rejection, ) +from .support.anima_nondiffusion_rejection_values import ( + synthetic_rejection_history, +) + def test_every_case_rejects_with_exact_issue_count_and_adapter_owner() -> None: """Require 60, one, two, and adapter-one aggregate issue boundaries.""" diff --git a/tests/test_anima_regional_output_capture.py b/tests/regional_generation/anima/test_anima_regional_output_capture.py similarity index 100% rename from tests/test_anima_regional_output_capture.py rename to tests/regional_generation/anima/test_anima_regional_output_capture.py diff --git a/tests/test_anima_regional_permutation_diagnostics.py b/tests/regional_generation/anima/test_anima_regional_permutation_diagnostics.py similarity index 98% rename from tests/test_anima_regional_permutation_diagnostics.py rename to tests/regional_generation/anima/test_anima_regional_permutation_diagnostics.py index d8f5bd4..d343765 100644 --- a/tests/test_anima_regional_permutation_diagnostics.py +++ b/tests/regional_generation/anima/test_anima_regional_permutation_diagnostics.py @@ -13,8 +13,6 @@ import comfy.ops import pytest import torch from comfy.ldm.anima.model import Anima -from regional_attention_test_values import single_entry_regions -from regional_lora_test_values import static_lora_schedule from simple_syrup.domain.regional_attention import RegionalAttentionBranch from simple_syrup.domain.regional_attention_batch import ( @@ -83,6 +81,13 @@ from simple_syrup.runtime.spatial_model_arguments import ( SPATIAL_BATCH_LAYOUT_KEY, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) +from ..regional.support.regional_lora_test_values import ( + static_lora_schedule, +) + @pytest.fixture(scope="module") def anima_surface() -> AnimaModuleSurface: diff --git a/tests/test_anima_regional_profile_probe.py b/tests/regional_generation/anima/test_anima_regional_profile_probe.py similarity index 100% rename from tests/test_anima_regional_profile_probe.py rename to tests/regional_generation/anima/test_anima_regional_profile_probe.py diff --git a/tests/test_anima_regional_self_attention.py b/tests/regional_generation/anima/test_anima_regional_self_attention.py similarity index 99% rename from tests/test_anima_regional_self_attention.py rename to tests/regional_generation/anima/test_anima_regional_self_attention.py index 54c5a44..113d4e1 100644 --- a/tests/test_anima_regional_self_attention.py +++ b/tests/regional_generation/anima/test_anima_regional_self_attention.py @@ -10,7 +10,6 @@ from typing import Any import pytest import torch -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -52,6 +51,10 @@ from simple_syrup.runtime.regional_self_attention_coherence import ( RegionalSelfAttentionCoherencePolicy, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _DeterministicSelfAttentionBase(nn.Module): """Provide shared deterministic installed-attention behavior.""" diff --git a/tests/test_anima_regression_oracle_artifact_validation.py b/tests/regional_generation/anima/test_anima_regression_oracle_artifact_validation.py similarity index 100% rename from tests/test_anima_regression_oracle_artifact_validation.py rename to tests/regional_generation/anima/test_anima_regression_oracle_artifact_validation.py diff --git a/tests/test_anima_regression_oracle_cli.py b/tests/regional_generation/anima/test_anima_regression_oracle_cli.py similarity index 93% rename from tests/test_anima_regression_oracle_cli.py rename to tests/regional_generation/anima/test_anima_regression_oracle_cli.py index ffc892b..298c8ae 100644 --- a/tests/test_anima_regression_oracle_cli.py +++ b/tests/regional_generation/anima/test_anima_regression_oracle_cli.py @@ -12,6 +12,8 @@ import subprocess import sys from pathlib import Path +from support.repository import REPOSITORY_ROOT + import tools.run_anima_regression_oracle as oracle from tools.anima_regression_oracle.manifest import OracleCommand, default_manifest @@ -19,7 +21,7 @@ from tools.anima_regression_oracle.manifest import OracleCommand, default_manife def test_oracle_script_supports_direct_filename_execution(tmp_path: Path) -> None: """Load repository-owned packages even when launched outside the repository.""" - script = Path(__file__).parents[1] / "tools" / "run_anima_regression_oracle.py" + script = REPOSITORY_ROOT / "tools" / "run_anima_regression_oracle.py" result = subprocess.run( [sys.executable, str(script), "--help"], @@ -53,7 +55,7 @@ def test_complete_level_includes_every_declared_managed_rerun() -> None: def test_every_managed_python_module_is_import_resolvable() -> None: """Reject stale or partially renamed managed command modules.""" - manifest = default_manifest(Path(__file__).parents[1]) + manifest = default_manifest(REPOSITORY_ROOT) modules = tuple( command.arguments[command.arguments.index("-m") + 1] for command in manifest.managed_rerun_commands diff --git a/tests/test_anima_regression_oracle_execution.py b/tests/regional_generation/anima/test_anima_regression_oracle_execution.py similarity index 100% rename from tests/test_anima_regression_oracle_execution.py rename to tests/regional_generation/anima/test_anima_regression_oracle_execution.py diff --git a/tests/test_anima_regression_oracle_manifest.py b/tests/regional_generation/anima/test_anima_regression_oracle_manifest.py similarity index 82% rename from tests/test_anima_regression_oracle_manifest.py rename to tests/regional_generation/anima/test_anima_regression_oracle_manifest.py index 47aeacf..e548e6f 100644 --- a/tests/test_anima_regression_oracle_manifest.py +++ b/tests/regional_generation/anima/test_anima_regression_oracle_manifest.py @@ -69,12 +69,21 @@ def test_manifest_names_focused_tests_explicitly() -> None: tests = tuple(argument for argument in arguments if argument.startswith("tests/")) assert len(tests) == 32 - assert "tests/test_anima_module_surface.py" in tests - assert "tests/test_anima_multi_lora_fidelity.py" in tests - assert "tests/test_anima_composition_phase.py" in tests - assert "tests/test_anima_tiled_attention_coupling_integration.py" in tests - assert "tests/test_anima_contextual_attention_coupling_integration.py" in tests - assert "tests/test_anima_regional_lora_performance_results.py" in tests + assert "tests/regional_generation/anima/test_anima_module_surface.py" in tests + assert "tests/regional_generation/anima/test_anima_multi_lora_fidelity.py" in tests + assert "tests/regional_generation/anima/test_anima_composition_phase.py" in tests + assert ( + "tests/regional_generation/anima/test_anima_tiled_attention_coupling_integration.py" + in tests + ) + assert ( + "tests/regional_generation/anima/test_anima_contextual_attention_coupling_integration.py" + in tests + ) + assert ( + "tests/regional_generation/anima/test_anima_regional_lora_performance_results.py" + in tests + ) def test_performance_command_supplies_an_external_artifact_inventory() -> None: diff --git a/tests/test_anima_regression_oracle_model_visibility.py b/tests/regional_generation/anima/test_anima_regression_oracle_model_visibility.py similarity index 100% rename from tests/test_anima_regression_oracle_model_visibility.py rename to tests/regional_generation/anima/test_anima_regression_oracle_model_visibility.py diff --git a/tests/test_anima_regression_oracle_results.py b/tests/regional_generation/anima/test_anima_regression_oracle_results.py similarity index 100% rename from tests/test_anima_regression_oracle_results.py rename to tests/regional_generation/anima/test_anima_regression_oracle_results.py diff --git a/tests/test_anima_self_attention_coherence.py b/tests/regional_generation/anima/test_anima_self_attention_coherence.py similarity index 100% rename from tests/test_anima_self_attention_coherence.py rename to tests/regional_generation/anima/test_anima_self_attention_coherence.py diff --git a/tests/test_anima_self_attention_ownership.py b/tests/regional_generation/anima/test_anima_self_attention_ownership.py similarity index 100% rename from tests/test_anima_self_attention_ownership.py rename to tests/regional_generation/anima/test_anima_self_attention_ownership.py diff --git a/tests/test_anima_self_attention_partition.py b/tests/regional_generation/anima/test_anima_self_attention_partition.py similarity index 100% rename from tests/test_anima_self_attention_partition.py rename to tests/regional_generation/anima/test_anima_self_attention_partition.py diff --git a/tests/test_anima_self_attention_work.py b/tests/regional_generation/anima/test_anima_self_attention_work.py similarity index 100% rename from tests/test_anima_self_attention_work.py rename to tests/regional_generation/anima/test_anima_self_attention_work.py diff --git a/tests/test_anima_single_adapter_mutations.py b/tests/regional_generation/anima/test_anima_single_adapter_mutations.py similarity index 99% rename from tests/test_anima_single_adapter_mutations.py rename to tests/regional_generation/anima/test_anima_single_adapter_mutations.py index fd23cf5..a2a1e44 100644 --- a/tests/test_anima_single_adapter_mutations.py +++ b/tests/regional_generation/anima/test_anima_single_adapter_mutations.py @@ -14,7 +14,6 @@ import torch from comfy.ldm.anima.model import Anima from comfy.ldm.cosmos.predict2 import Attention, Block, GPT2FeedForward from comfy.patcher_extension import CallbacksMP -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -99,6 +98,10 @@ from simple_syrup.runtime.regional_lora.execution_cache import ( ) from simple_syrup.runtime.regional_lora.standard_adapter import StandardLoraTarget +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _AnimaModelRoot(nn.Module): """Expose installed Anima under Comfy's object-patch root path.""" diff --git a/tests/test_anima_tiled_accumulation.py b/tests/regional_generation/anima/test_anima_tiled_accumulation.py similarity index 99% rename from tests/test_anima_tiled_accumulation.py rename to tests/regional_generation/anima/test_anima_tiled_accumulation.py index 19e3c17..778c750 100644 --- a/tests/test_anima_tiled_accumulation.py +++ b/tests/regional_generation/anima/test_anima_tiled_accumulation.py @@ -10,7 +10,6 @@ from typing import Any, cast import pytest import torch -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -62,6 +61,10 @@ from simple_syrup.runtime.tile_prediction_accumulation import ( TilePredictionAccumulator, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _ContextValueAttention(nn.Module): """Return each compact row's context scalar and record branch identity.""" diff --git a/tests/test_anima_tiled_attention_coupling_integration.py b/tests/regional_generation/anima/test_anima_tiled_attention_coupling_integration.py similarity index 100% rename from tests/test_anima_tiled_attention_coupling_integration.py rename to tests/regional_generation/anima/test_anima_tiled_attention_coupling_integration.py diff --git a/tests/test_anima_tiled_branch_pruning.py b/tests/regional_generation/anima/test_anima_tiled_branch_pruning.py similarity index 98% rename from tests/test_anima_tiled_branch_pruning.py rename to tests/regional_generation/anima/test_anima_tiled_branch_pruning.py index 8033f66..1505018 100644 --- a/tests/test_anima_tiled_branch_pruning.py +++ b/tests/regional_generation/anima/test_anima_tiled_branch_pruning.py @@ -7,7 +7,6 @@ from __future__ import annotations import torch -from regional_attention_test_values import single_entry_regions from torch import nn from simple_syrup.domain.regional_attention import RegionalAttentionBranch @@ -55,6 +54,10 @@ from simple_syrup.runtime.regional_lora.anima_query_masks import ( AnimaQueryMaskProjector, ) +from ..regional.support.regional_attention_test_values import ( + single_entry_regions, +) + class _ContextValueAttention(nn.Module): """Return one context scalar per row and record compact branch identity.""" diff --git a/tests/test_anima_visual_performance_workflow.py b/tests/regional_generation/anima/test_anima_visual_performance_workflow.py similarity index 100% rename from tests/test_anima_visual_performance_workflow.py rename to tests/regional_generation/anima/test_anima_visual_performance_workflow.py diff --git a/tests/regional_generation/attention_coupling/__init__.py b/tests/regional_generation/attention_coupling/__init__.py new file mode 100644 index 0000000..79d3dd6 --- /dev/null +++ b/tests/regional_generation/attention_coupling/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation attention coupling test behavior.""" diff --git a/tests/regional_generation/attention_coupling/support/__init__.py b/tests/regional_generation/attention_coupling/support/__init__.py new file mode 100644 index 0000000..3bd8ea3 --- /dev/null +++ b/tests/regional_generation/attention_coupling/support/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation attention coupling support test behavior.""" diff --git a/tests/attention_coupling_diagnostics_harness.py b/tests/regional_generation/attention_coupling/support/attention_coupling_diagnostics_harness.py similarity index 95% rename from tests/attention_coupling_diagnostics_harness.py rename to tests/regional_generation/attention_coupling/support/attention_coupling_diagnostics_harness.py index 03c6e5f..2b534ed 100644 --- a/tests/attention_coupling_diagnostics_harness.py +++ b/tests/regional_generation/attention_coupling/support/attention_coupling_diagnostics_harness.py @@ -7,10 +7,6 @@ from __future__ import annotations import torch -from attention_coupling_invariant_contract import ( - AttentionCouplingDiagnosticsObservation, -) -from attention_coupling_invariant_values import scheduled_invariant_plan from simple_syrup.domain.spatial_views import ( SpatialBatchLayout, @@ -27,6 +23,13 @@ from simple_syrup.runtime.regional_attention_model_call import ( RegionalAttentionModelCallResolver, ) +from .attention_coupling_invariant_contract import ( + AttentionCouplingDiagnosticsObservation, +) +from .attention_coupling_invariant_values import ( + scheduled_invariant_plan, +) + def build_diagnostics_observation( runtime_type: type[object], diff --git a/tests/attention_coupling_invariant_contract.py b/tests/regional_generation/attention_coupling/support/attention_coupling_invariant_contract.py similarity index 100% rename from tests/attention_coupling_invariant_contract.py rename to tests/regional_generation/attention_coupling/support/attention_coupling_invariant_contract.py diff --git a/tests/attention_coupling_invariant_values.py b/tests/regional_generation/attention_coupling/support/attention_coupling_invariant_values.py similarity index 100% rename from tests/attention_coupling_invariant_values.py rename to tests/regional_generation/attention_coupling/support/attention_coupling_invariant_values.py diff --git a/tests/test_attention_coupling_backend_invariants.py b/tests/regional_generation/attention_coupling/test_attention_coupling_backend_invariants.py similarity index 91% rename from tests/test_attention_coupling_backend_invariants.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_backend_invariants.py index a24b0a9..8988778 100644 --- a/tests/test_attention_coupling_backend_invariants.py +++ b/tests/regional_generation/attention_coupling/test_attention_coupling_backend_invariants.py @@ -10,17 +10,34 @@ from typing import cast import pytest import torch -from anima_attention_coupling_call_harness import AnimaAttentionCouplingCallHarness -from anima_attention_coupling_diagnostics_harness import ( + +from simple_syrup.domain.regional_attention import RegionalAttentionBranch + +from ..anima.support.anima_attention_coupling_call_harness import ( + AnimaAttentionCouplingCallHarness, +) +from ..anima.support.anima_attention_coupling_diagnostics_harness import ( AnimaAttentionCouplingDiagnosticsHarness, ) -from anima_attention_coupling_invariant_harness import ( +from ..anima.support.anima_attention_coupling_invariant_harness import ( AnimaAttentionCouplingInvariantHarness, ) -from anima_attention_coupling_lifecycle_harness import ( +from ..anima.support.anima_attention_coupling_lifecycle_harness import ( AnimaAttentionCouplingLifecycleHarness, ) -from attention_coupling_invariant_contract import ( +from ..unet.support.unet_attention_coupling_call_harness import ( + UnetAttentionCouplingCallHarness, +) +from ..unet.support.unet_attention_coupling_diagnostics_harness import ( + UnetAttentionCouplingDiagnosticsHarness, +) +from ..unet.support.unet_attention_coupling_invariant_harness import ( + UnetAttentionCouplingInvariantHarness, +) +from ..unet.support.unet_attention_coupling_lifecycle_harness import ( + UnetAttentionCouplingLifecycleHarness, +) +from .support.attention_coupling_invariant_contract import ( AttentionCouplingCallHarness, AttentionCouplingCallScenario, AttentionCouplingDiagnosticsHarness, @@ -29,19 +46,9 @@ from attention_coupling_invariant_contract import ( AttentionCouplingInvariantScenario, AttentionCouplingLifecycleHarness, ) -from attention_coupling_invariant_values import scheduled_invariant_plan -from unet_attention_coupling_call_harness import UnetAttentionCouplingCallHarness -from unet_attention_coupling_diagnostics_harness import ( - UnetAttentionCouplingDiagnosticsHarness, +from .support.attention_coupling_invariant_values import ( + scheduled_invariant_plan, ) -from unet_attention_coupling_invariant_harness import ( - UnetAttentionCouplingInvariantHarness, -) -from unet_attention_coupling_lifecycle_harness import ( - UnetAttentionCouplingLifecycleHarness, -) - -from simple_syrup.domain.regional_attention import RegionalAttentionBranch def _single_entries( diff --git a/tests/test_attention_coupling_benchmark_manifest.py b/tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_manifest.py similarity index 100% rename from tests/test_attention_coupling_benchmark_manifest.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_manifest.py diff --git a/tests/test_attention_coupling_benchmark_masks.py b/tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_masks.py similarity index 100% rename from tests/test_attention_coupling_benchmark_masks.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_masks.py diff --git a/tests/test_attention_coupling_benchmark_probe.py b/tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_probe.py similarity index 100% rename from tests/test_attention_coupling_benchmark_probe.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_probe.py diff --git a/tests/test_attention_coupling_benchmark_results.py b/tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_results.py similarity index 100% rename from tests/test_attention_coupling_benchmark_results.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_results.py diff --git a/tests/test_attention_coupling_benchmark_workflow.py b/tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_workflow.py similarity index 100% rename from tests/test_attention_coupling_benchmark_workflow.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_benchmark_workflow.py diff --git a/tests/test_attention_coupling_bypass_integration.py b/tests/regional_generation/attention_coupling/test_attention_coupling_bypass_integration.py similarity index 100% rename from tests/test_attention_coupling_bypass_integration.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_bypass_integration.py diff --git a/tests/test_attention_coupling_interop_preflight.py b/tests/regional_generation/attention_coupling/test_attention_coupling_interop_preflight.py similarity index 100% rename from tests/test_attention_coupling_interop_preflight.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_interop_preflight.py diff --git a/tests/test_attention_coupling_model_family_selector.py b/tests/regional_generation/attention_coupling/test_attention_coupling_model_family_selector.py similarity index 100% rename from tests/test_attention_coupling_model_family_selector.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_model_family_selector.py diff --git a/tests/test_attention_coupling_model_preparation_routing.py b/tests/regional_generation/attention_coupling/test_attention_coupling_model_preparation_routing.py similarity index 98% rename from tests/test_attention_coupling_model_preparation_routing.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_model_preparation_routing.py index 5019412..9090ba9 100644 --- a/tests/test_attention_coupling_model_preparation_routing.py +++ b/tests/regional_generation/attention_coupling/test_attention_coupling_model_preparation_routing.py @@ -15,17 +15,10 @@ import comfy.model_patcher import pytest import torch from comfy.ldm.anima.model import Anima as AnimaDiffusionModel -from regional_model_capability_test_values import ( - AlternateImageLatent, - standard_unet_graph, -) -from regional_model_capability_test_values import ( - empty_module as _empty_module, -) -from regional_model_capability_test_values import ( - patcher as _patcher, -) +from simple_syrup.domain.attention_coupling_preparation import ( + AttentionCouplingPreparation, +) from simple_syrup.domain.conditioning_batch import ConditioningBatch from simple_syrup.domain.conditioning_schedule import ConditioningScheduleRange from simple_syrup.domain.processed_regional_attention import ( @@ -51,9 +44,6 @@ from simple_syrup.services.anima_attention_coupling_model_family import ( from simple_syrup.services.attention_coupling_model_preparation_service import ( AttentionCouplingModelPreparationService, ) -from simple_syrup.services.attention_coupling_preparation_service import ( - AttentionCouplingPreparation, -) from simple_syrup.services.attention_coupling_sampling_service import ( AttentionCouplingSamplingService, ) @@ -64,6 +54,17 @@ from simple_syrup.services.tiled_attention_coupling_sampling_service import ( TiledAttentionCouplingSamplingService, ) +from ..regional.support.regional_model_capability_test_values import ( + AlternateImageLatent, + standard_unet_graph, +) +from ..regional.support.regional_model_capability_test_values import ( + empty_module as _empty_module, +) +from ..regional.support.regional_model_capability_test_values import ( + patcher as _patcher, +) + class _LoraAdapter: """Return empty regional model-adapter state for routing tests.""" diff --git a/tests/test_attention_coupling_model_preparation_service.py b/tests/regional_generation/attention_coupling/test_attention_coupling_model_preparation_service.py similarity index 99% rename from tests/test_attention_coupling_model_preparation_service.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_model_preparation_service.py index cc619d0..ea3dba5 100644 --- a/tests/test_attention_coupling_model_preparation_service.py +++ b/tests/regional_generation/attention_coupling/test_attention_coupling_model_preparation_service.py @@ -13,6 +13,9 @@ from uuid import uuid4 import pytest import torch +from simple_syrup.domain.attention_coupling_preparation import ( + AttentionCouplingPreparation, +) from simple_syrup.domain.conditioning_batch import ConditioningBatch from simple_syrup.domain.conditioning_schedule import ConditioningScheduleRange from simple_syrup.domain.processed_regional_attention import ( @@ -40,9 +43,6 @@ from simple_syrup.services.attention_coupling_model_family import ( from simple_syrup.services.attention_coupling_model_preparation_service import ( AttentionCouplingModelPreparationService, ) -from simple_syrup.services.attention_coupling_preparation_service import ( - AttentionCouplingPreparation, -) _PREPARATION_EVENTS: list[str] = [] diff --git a/tests/test_attention_coupling_phase_node.py b/tests/regional_generation/attention_coupling/test_attention_coupling_phase_node.py similarity index 100% rename from tests/test_attention_coupling_phase_node.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_phase_node.py diff --git a/tests/test_attention_coupling_phase_profile.py b/tests/regional_generation/attention_coupling/test_attention_coupling_phase_profile.py similarity index 100% rename from tests/test_attention_coupling_phase_profile.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_phase_profile.py diff --git a/tests/test_attention_coupling_preparation_service.py b/tests/regional_generation/attention_coupling/test_attention_coupling_preparation_service.py similarity index 100% rename from tests/test_attention_coupling_preparation_service.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_preparation_service.py diff --git a/tests/test_attention_coupling_prepared_model_cache.py b/tests/regional_generation/attention_coupling/test_attention_coupling_prepared_model_cache.py similarity index 100% rename from tests/test_attention_coupling_prepared_model_cache.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_prepared_model_cache.py diff --git a/tests/test_attention_coupling_request.py b/tests/regional_generation/attention_coupling/test_attention_coupling_request.py similarity index 100% rename from tests/test_attention_coupling_request.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_request.py diff --git a/tests/test_attention_coupling_sampling_service.py b/tests/regional_generation/attention_coupling/test_attention_coupling_sampling_service.py similarity index 100% rename from tests/test_attention_coupling_sampling_service.py rename to tests/regional_generation/attention_coupling/test_attention_coupling_sampling_service.py diff --git a/tests/test_contextual_attention_coupling_sampling_service.py b/tests/regional_generation/attention_coupling/test_contextual_attention_coupling_sampling_service.py similarity index 100% rename from tests/test_contextual_attention_coupling_sampling_service.py rename to tests/regional_generation/attention_coupling/test_contextual_attention_coupling_sampling_service.py diff --git a/tests/test_tiled_attention_coupling_sampling_service.py b/tests/regional_generation/attention_coupling/test_tiled_attention_coupling_sampling_service.py similarity index 100% rename from tests/test_tiled_attention_coupling_sampling_service.py rename to tests/regional_generation/attention_coupling/test_tiled_attention_coupling_sampling_service.py diff --git a/tests/regional_generation/attention_regions/__init__.py b/tests/regional_generation/attention_regions/__init__.py new file mode 100644 index 0000000..0d42181 --- /dev/null +++ b/tests/regional_generation/attention_regions/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation attention regions test behavior.""" diff --git a/tests/test_attention_region_cadence.py b/tests/regional_generation/attention_regions/test_attention_region_cadence.py similarity index 100% rename from tests/test_attention_region_cadence.py rename to tests/regional_generation/attention_regions/test_attention_region_cadence.py diff --git a/tests/regional_generation/attention_regions/test_attention_region_capture.py b/tests/regional_generation/attention_regions/test_attention_region_capture.py new file mode 100644 index 0000000..2600cf5 --- /dev/null +++ b/tests/regional_generation/attention_regions/test_attention_region_capture.py @@ -0,0 +1,439 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Test observation-only selected-token attention capture.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.domain.attention_region_capture import ( + AttentionCapturePlan, + AttentionCaptureProfile, + AttentionRegionControls, + AttentionRegionRequest, + AttentionRegionRequestKind, +) +from simple_syrup.domain.attention_region_maps import ( + AttentionTokenCatalog, + AttentionTokenSpan, +) +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily +from simple_syrup.runtime.attention_region_affinity import ( + _select_relevant_head_maps, +) +from simple_syrup.runtime.attention_region_capture import AttentionRegionCaptureSession +from simple_syrup.runtime.attention_region_contextual_spans import ( + ATTENTION_CONTEXTUAL_SPAN_SELECTOR, +) +from simple_syrup.runtime.attention_region_phrase_evidence import ( + ANIMA_PHRASE_EVIDENCE_SERVICE, +) + + +def test_attention_controls_reject_invalid_instance_recall() -> None: + """Reject disconnected-instance recall outside its normalized range.""" + + with pytest.raises(ValueError, match="instance recall"): + AttentionRegionControls( + 0.0, + 1.0, + 0.15, + 0.25, + 0.0, + 1, + AttentionCaptureProfile.FAST, + instance_recall=1.01, + ) + + +def test_attention_controls_reject_invalid_geometry_recall() -> None: + """Reject connected-geometry recall outside its normalized range.""" + + with pytest.raises(ValueError, match="geometry recall"): + AttentionRegionControls( + 0.0, + 1.0, + 0.15, + 0.25, + 0.0, + 1, + AttentionCaptureProfile.FAST, + geometry_recall=-0.01, + ) + + +def test_capture_selects_positive_rows_and_exact_prompt_tokens() -> None: + """Capture one native concept without retaining a full attention matrix.""" + + session = _session(profile=AttentionCaptureProfile.EXHAUSTIVE) + query = torch.tensor( + [ + [[1.0, 0.0], [0.0, 1.0]], + [[-1.0, 0.0], [0.0, -1.0]], + ] + ) + key = torch.tensor( + [ + [[0.0, 0.0], [1.0, 0.0], [0.0, 1.0]], + [[0.0, 0.0], [-1.0, 0.0], [0.0, -1.0]], + ] + ) + + session.observe( + query, + key, + key, + 1, + _options(cond_or_uncond=[0, 1]), + skip_reshape=False, + ) + + maps = session.maps_for("search") + assert len(maps) == 1 + assert maps[0].label == "pink hair" + assert maps[0].values.device.type == "cpu" + assert maps[0].values[0].item() > maps[0].values[1].item() + + +def test_capture_retains_value_aware_phrase_evidence_beside_raw_attention() -> None: + """Favor the phrase token carrying stronger projected model contribution.""" + + session = _session( + profile=AttentionCaptureProfile.EXHAUSTIVE, + sequence_length=4, + token_indices=(1, 2), + ) + query = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]]) + key = torch.tensor([[[0.0, 0.0], [2.0, 0.0], [0.0, 2.0], [0.0, 0.0]]]) + value = torch.tensor([[[1.0, 1.0], [0.1, 0.1], [4.0, 4.0], [1.0, 1.0]]]) + + session.observe( + query, + key, + value, + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + attention_map = session.maps_for("search")[0] + assert torch.isclose(attention_map.values[0], attention_map.values[1]) + assert attention_map.concept_values is not None + assert attention_map.concept_values[1] > attention_map.concept_values[0] + assert attention_map.uniform_probability == 0.25 + + +def test_contextual_span_selector_admits_only_distinct_related_prompt_spans() -> None: + """Select a related conditioning span while rejecting an orthogonal concept.""" + + target = AttentionTokenSpan("pink hair", 1, (1,)) + related = AttentionTokenSpan("twintails", 1, (2,)) + unrelated = AttentionTokenSpan("smile", 1, (3,)) + key = torch.tensor([[[[0.0, 0.0], [1.0, 0.0], [0.9, 0.1], [0.0, 1.0]]]]) + + selections = ATTENTION_CONTEXTUAL_SPAN_SELECTOR.select( + key=key, + targets=(target,), + candidates=(target, related, unrelated), + ) + + assert selections[target].token_indices == (1, 2) + assert selections[target].token_weights[0] == 1.0 + assert 0.0 < selections[target].token_weights[1] < 1.0 + + +def test_contextual_span_selector_uses_object_head_instead_of_modifier_color() -> None: + """Relate an object phrase through its head without following a color token.""" + + target = AttentionTokenSpan("pink outfit", 1, (1, 2), (2,)) + related_part = AttentionTokenSpan("short skirt", 1, (3, 4), (4,)) + color_only = AttentionTokenSpan("pink petals", 1, (5,), (5,)) + key = torch.tensor( + [ + [ + [ + [0.0, 0.0], + [0.0, 1.0], + [1.0, 0.0], + [0.0, 1.0], + [0.9, 0.1], + [0.0, 1.0], + ] + ] + ] + ) + + selection = ATTENTION_CONTEXTUAL_SPAN_SELECTOR.select( + key=key, + targets=(target,), + candidates=(target, related_part, color_only), + )[target] + + assert selection.token_indices == (1, 2, 4) + assert selection.token_weights[0] < selection.token_weights[1] + + +def test_anima_compound_phrase_uses_specific_modifier_to_constrain_its_head() -> None: + """Keep a small compound concept local when its noun head is spatially broad.""" + + probability = torch.tensor( + [ + [ + [ + [0.8, 1.0, 1.0], + [0.8, 0.9, 0.9], + [0.8, 0.05, 0.7], + [0.8, 0.05, 0.7], + ] + ] + ] + ) + span = AttentionTokenSpan( + "blue butterfly ornaments", + 1, + (0, 1, 2), + (2,), + ) + + concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( + probability=probability, + span=span, + union_positions={0: 0, 1: 1, 2: 2}, + ) + + assert concept[0, 0] > 0.8 + assert concept[0, 1] > 0.7 + assert concept[0, 2] < concept[0, 0] * 0.5 + assert concept[0, 3] < concept[0, 0] * 0.5 + + +def test_anima_broad_modifier_does_not_erase_extended_head_geometry() -> None: + """Preserve a noun silhouette when its only modifier carries no locality.""" + + probability = torch.tensor( + [ + [ + [ + [0.6, 1.0], + [0.6, 0.8], + [0.6, 0.5], + [0.6, 0.3], + ] + ] + ] + ) + span = AttentionTokenSpan("pink hair", 1, (0, 1), (1,)) + + concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( + probability=probability, + span=span, + union_positions={0: 0, 1: 1}, + ) + + assert torch.allclose(concept[0], probability[0, 0, :, 1].to(torch.float16)) + + +def test_anima_single_token_concept_uses_noun_head_evidence_directly() -> None: + """Handle noun-only prompt segments without requiring modifier positions.""" + + probability = torch.tensor([[[[1.0], [0.8], [0.3], [0.0]]]]) + span = AttentionTokenSpan("twintails", 1, (0,), (0,)) + + concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( + probability=probability, + span=span, + union_positions={0: 0}, + ) + + assert torch.equal(concept[0], probability[0, 0, :, 0].to(torch.float16)) + + +def test_anima_related_prompt_head_recovers_a_disjoint_concept_part() -> None: + """Admit a related prompt noun at reduced strength without replacing the core.""" + + probability = torch.tensor( + [ + [ + [ + [0.5, 1.0, 0.0], + [0.5, 0.2, 0.0], + [0.5, 0.0, 0.9], + [0.5, 0.0, 0.8], + ] + ] + ] + ) + span = AttentionTokenSpan("pink hair", 1, (0, 1), (1,)) + + concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( + probability=probability, + span=span, + union_positions={0: 0, 1: 1, 2: 2}, + contextual_token_indices=(0, 1, 2), + contextual_token_weights=(0.45, 1.0, 0.3), + ) + + assert concept[0, 0] == 1.0 + assert concept[0, 2] > 0.25 + assert concept[0, 3] > 0.2 + + +def test_concept_head_selection_rejects_a_spatially_disagreeing_head() -> None: + """Prefer concept heads that agree on the object while retaining their detail.""" + + probability = torch.tensor( + [ + [ + [[0.9], [0.8], [0.4], [0.0]], + [[0.8], [0.9], [0.0], [0.0]], + [[0.7], [0.8], [0.0], [0.0]], + [[0.0], [0.0], [0.9], [0.9]], + ] + ] + ) + + _weighted, selected = _select_relevant_head_maps( + probability, + torch.ones(1, 4, 1), + ) + + assert selected[0, 0, 0] > selected[0, 3, 0] + assert selected[0, 1, 0] > selected[0, 3, 0] + + +def test_related_tokens_use_heads_selected_by_the_exact_concept() -> None: + """Recover related detail without admitting its independently face-focused head.""" + + probability = torch.tensor( + [ + [ + [[0.9, 0.1], [0.8, 0.1], [0.4, 0.8], [0.0, 0.0]], + [[0.8, 0.1], [0.9, 0.1], [0.4, 0.7], [0.0, 0.0]], + [[0.7, 0.1], [0.8, 0.1], [0.3, 0.6], [0.0, 0.0]], + [[0.0, 0.0], [0.0, 0.0], [0.0, 0.1], [0.9, 0.9]], + ] + ] + ) + + _weighted, selected = _select_relevant_head_maps( + probability, + torch.ones(1, 4, 2), + exact_token_mask=torch.tensor([True, False]), + ) + + assert selected[0, 2, 1] > selected[0, 3, 1] + + +def test_concept_capture_enriches_related_span_without_changing_raw_attention() -> None: + """Recover self-grouped related evidence only in the derived concept channel.""" + + target = AttentionTokenSpan("pink hair", 1, (1,)) + related = AttentionTokenSpan("twintails", 1, (2,)) + unrelated = AttentionTokenSpan("smile", 1, (3,)) + session = _session( + profile=AttentionCaptureProfile.EXHAUSTIVE, + sequence_length=4, + catalog_spans=(target, related, unrelated), + ) + query = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]]) + key = torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.6, 0.8], [0.0, -1.0]]]) + value = torch.ones_like(key) + + spatial = torch.tensor([[[1.0, 0.0], [1.0, 0.0]]]) + session.observe( + spatial, + spatial, + spatial, + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + session.observe( + query, + key, + value, + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + attention_map = session.maps_for("search")[0] + assert attention_map.concept_values is not None + raw_ratio = attention_map.values[1] / attention_map.values[0] + concept_ratio = attention_map.concept_values[1] / attention_map.concept_values[0] + assert concept_ratio > raw_ratio + + +def _session( + profile: AttentionCaptureProfile = AttentionCaptureProfile.EXHAUSTIVE, + sequence_length: int = 3, + source_aspect: float | None = None, + model_family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, + token_indices: tuple[int, ...] = (1,), + catalog_spans: tuple[AttentionTokenSpan, ...] | None = None, +) -> AttentionRegionCaptureSession: + """Return one exact-native-query capture session.""" + + controls = AttentionRegionControls(0.0, 1.0, 0.3, 0.2, 0.5, 1, profile) + request = AttentionRegionRequest( + "search", + AttentionRegionRequestKind.CONCEPT_SEGS, + ("pink hair",), + controls, + ) + plan = AttentionCapturePlan( + "sampler", + "sampler", + "model", + ("model", 0), + ("positive", 0), + (request,), + "pink hair", + ("loader", 1), + source_aspect, + ) + default_span = AttentionTokenSpan("pink hair", 1, token_indices) + catalog = AttentionTokenCatalog( + sequence_length, + catalog_spans or (default_span,), + tuple(range(sequence_length)), + ) + return AttentionRegionCaptureSession( + plan=plan, + model_family=model_family, + token_catalog=catalog, + request_spans={"search": (catalog.spans[0],)}, + ) + + +def _options( + *, + cond_or_uncond: list[int], + sigma: float = 1.0, + block_index: int = 0, +) -> dict[str, object]: + """Return exact sampler metadata at the beginning of denoising.""" + + return { + "sample_sigmas": torch.tensor([1.0, 0.5, 0.0]), + "sigmas": torch.tensor([sigma]), + "cond_or_uncond": cond_or_uncond, + "block": ("middle", 0), + "block_index": block_index, + } + + +def _patcher() -> Any: + """Create a real CPU ModelPatcher with isolated transformer options.""" + + from comfy.model_patcher import ModelPatcher + + base_model = torch.nn.Module() + base_model.diffusion_model = torch.nn.Linear(1, 1) + device = torch.device("cpu") + return ModelPatcher(base_model, load_device=device, offload_device=device) diff --git a/tests/regional_generation/attention_regions/test_attention_region_capture_backend.py b/tests/regional_generation/attention_regions/test_attention_region_capture_backend.py new file mode 100644 index 0000000..b81eccd --- /dev/null +++ b/tests/regional_generation/attention_regions/test_attention_region_capture_backend.py @@ -0,0 +1,310 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Test observation-only selected-token attention capture.""" + +from __future__ import annotations + +from typing import Any + +import torch + +from simple_syrup.domain.attention_region_capture import ( + AttentionCapturePlan, + AttentionCaptureProfile, + AttentionRegionControls, + AttentionRegionRequest, + AttentionRegionRequestKind, +) +from simple_syrup.domain.attention_region_maps import ( + AttentionTokenCatalog, + AttentionTokenSpan, + OpenVocabularyContext, +) +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily +from simple_syrup.runtime.attention_region_capture import AttentionRegionCaptureSession +from simple_syrup.runtime.attention_region_capture_backend import ( + ATTENTION_REGION_CAPTURE_BACKEND, + OptimizedAttentionCaptureOverride, +) +from simple_syrup.runtime.attention_region_open_vocabulary import ( + OpenVocabularyKeyProjector, +) + + +def test_backend_clones_model_and_composes_existing_attention_override() -> None: + """Keep source options intact while preserving an upstream override.""" + + source = _patcher() + + def previous( + original: object, + *args: object, + **kwargs: object, + ) -> torch.Tensor: + """Return a sentinel from an admitted upstream override.""" + + del original, args, kwargs + return torch.tensor(3.0) + + source.model_options["transformer_options"]["optimized_attention_override"] = ( + previous + ) + + derived: Any = ATTENTION_REGION_CAPTURE_BACKEND.derive(source, _session()) + + assert derived.parent is source + assert ( + source.model_options["transformer_options"]["optimized_attention_override"] + is previous + ) + installed = derived.model_options["transformer_options"][ + "optimized_attention_override" + ] + assert isinstance(installed, OptimizedAttentionCaptureOverride) + assert ( + installed( + lambda *_args, **_kwargs: torch.tensor(1.0), + torch.ones(1, 1, 1), + torch.ones(1, 1, 1), + torch.ones(1, 1, 1), + 1, + ).item() + == 3.0 + ) + + +def test_open_vocabulary_query_projects_side_keys_without_changing_native_keys() -> ( + None +): + """Reuse SDXL spatial queries for an absent phrase without conditioning edits.""" + + controls = AttentionRegionControls( + 0.0, + 1.0, + 0.3, + 0.2, + 0.5, + 1, + AttentionCaptureProfile.EXHAUSTIVE, + ) + request = AttentionRegionRequest( + "search", + AttentionRegionRequestKind.CONCEPT_SEGS, + ("head",), + controls, + ) + plan = AttentionCapturePlan( + "sampler", + "sampler", + "model", + ("model", 0), + ("positive", 0), + (request,), + "1girl", + ("loader", 1), + ) + context = OpenVocabularyContext( + "head", + torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 0.0]]]), + (1,), + ) + session = AttentionRegionCaptureSession( + plan=plan, + model_family=RegionalModelFamily.STANDARD_UNET, + token_catalog=AttentionTokenCatalog( + 3, (AttentionTokenSpan("1girl", 1, (1,)),), (1, 2, 3) + ), + request_spans={"search": ()}, + open_vocabulary_contexts=(context,), + ) + projector = OpenVocabularyKeyProjector(torch.nn.Identity(), (context,), session) + native_context = torch.tensor([[[2.0, 0.0], [0.0, 2.0], [0.0, 0.0]]]) + + native_keys = projector(native_context) + session.observe( + torch.tensor([[[1.0, 0.0], [0.0, 1.0]]]), + native_keys, + native_context, + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + assert torch.equal(native_keys, native_context) + maps = session.maps_for("search") + assert len(maps) == 1 + assert maps[0].label == "head" + assert maps[0].values[0].item() > maps[0].values[1].item() + + +def test_open_vocabulary_query_projection_is_cached_per_device_and_dtype() -> None: + """Avoid repeating absent-query key projections at every denoising call.""" + + class CountingProjection(torch.nn.Module): + """Count native and side-projection calls while preserving values.""" + + calls: int + + def __init__(self) -> None: + """Initialize an unused projection counter.""" + + super().__init__() + self.calls = 0 + + def forward(self, value: torch.Tensor) -> torch.Tensor: + """Return the input after recording the projection.""" + + self.calls += 1 + return value + + context = OpenVocabularyContext( + "head", + torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 0.0]]]), + (1,), + ) + projection = CountingProjection() + projector = OpenVocabularyKeyProjector(projection, (context,), _session()) + native_context = torch.ones(1, 3, 2) + + assert torch.equal(projector(native_context), native_context) + assert torch.equal(projector(native_context), native_context) + assert projection.calls == 3 + + +def test_sampled_denominator_bounds_unusually_strong_selected_keys() -> None: + """Keep fast-profile maps finite when denominator sampling misses the peak key.""" + + session = _session( + profile=AttentionCaptureProfile.FAST, + sequence_length=40, + ) + query = torch.tensor([[[1000.0, 0.0], [0.0, 1.0]]]) + key = torch.zeros(1, 40, 2) + key[0, 1, 0] = 1000.0 + + session.observe( + query, + key, + key, + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + maps = session.maps_for("search") + assert len(maps) == 1 + assert torch.isfinite(maps[0].values).all().item() + assert maps[0].values.max().item() <= 1.0 + + +def test_graph_source_aspect_orients_dit_geometry_without_transformer_metadata() -> ( + None +): + """Use the selected sampler's portrait input when a DiT omits shape metadata.""" + + session = _session(source_aspect=0.75) + + session.observe( + torch.ones(1, 12, 2), + torch.ones(1, 3, 2), + torch.ones(1, 3, 2), + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + attention_map = session.maps_for("search")[0] + assert (attention_map.spatial_height, attention_map.spatial_width) == (4, 3) + + +def test_anima_without_graph_or_runtime_geometry_fails_closed() -> None: + """Return no maps instead of guessing an unprojectable Anima grid orientation.""" + + session = _session( + sequence_length=512, + model_family=RegionalModelFamily.ANIMA, + ) + + session.observe( + torch.ones(1, 12, 2), + torch.ones(1, 512, 2), + torch.ones(1, 512, 2), + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + assert session.maps_for("search") == () + assert "not graph-visible" in session.status_message + + +def _session( + profile: AttentionCaptureProfile = AttentionCaptureProfile.EXHAUSTIVE, + sequence_length: int = 3, + source_aspect: float | None = None, + model_family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, + token_indices: tuple[int, ...] = (1,), + catalog_spans: tuple[AttentionTokenSpan, ...] | None = None, +) -> AttentionRegionCaptureSession: + """Return one exact-native-query capture session.""" + + controls = AttentionRegionControls(0.0, 1.0, 0.3, 0.2, 0.5, 1, profile) + request = AttentionRegionRequest( + "search", + AttentionRegionRequestKind.CONCEPT_SEGS, + ("pink hair",), + controls, + ) + plan = AttentionCapturePlan( + "sampler", + "sampler", + "model", + ("model", 0), + ("positive", 0), + (request,), + "pink hair", + ("loader", 1), + source_aspect, + ) + default_span = AttentionTokenSpan("pink hair", 1, token_indices) + catalog = AttentionTokenCatalog( + sequence_length, + catalog_spans or (default_span,), + tuple(range(sequence_length)), + ) + return AttentionRegionCaptureSession( + plan=plan, + model_family=model_family, + token_catalog=catalog, + request_spans={"search": (catalog.spans[0],)}, + ) + + +def _options( + *, + cond_or_uncond: list[int], + sigma: float = 1.0, + block_index: int = 0, +) -> dict[str, object]: + """Return exact sampler metadata at the beginning of denoising.""" + + return { + "sample_sigmas": torch.tensor([1.0, 0.5, 0.0]), + "sigmas": torch.tensor([sigma]), + "cond_or_uncond": cond_or_uncond, + "block": ("middle", 0), + "block_index": block_index, + } + + +def _patcher() -> Any: + """Create a real CPU ModelPatcher with isolated transformer options.""" + + from comfy.model_patcher import ModelPatcher + + base_model = torch.nn.Module() + base_model.diffusion_model = torch.nn.Linear(1, 1) + device = torch.device("cpu") + return ModelPatcher(base_model, load_device=device, offload_device=device) diff --git a/tests/regional_generation/attention_regions/test_attention_region_completion.py b/tests/regional_generation/attention_regions/test_attention_region_completion.py new file mode 100644 index 0000000..3b7559a --- /dev/null +++ b/tests/regional_generation/attention_regions/test_attention_region_completion.py @@ -0,0 +1,503 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Test observation-only selected-token attention capture.""" + +from __future__ import annotations + +from concurrent.futures import ThreadPoolExecutor +from threading import Event +from typing import Any + +import torch + +from simple_syrup.domain.attention_region_capture import ( + AttentionCapturePlan, + AttentionCaptureProfile, + AttentionRegionControls, + AttentionRegionRequest, + AttentionRegionRequestKind, +) +from simple_syrup.domain.attention_region_maps import ( + AttentionTokenCatalog, + AttentionTokenSpan, +) +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily +from simple_syrup.runtime.attention_region_affinity import ( + ATTENTION_AFFINITY_CALCULATOR, +) +from simple_syrup.runtime.attention_region_capture import AttentionRegionCaptureSession +from simple_syrup.runtime.attention_region_capture_backend import ( + OptimizedAttentionCaptureOverride, +) +from simple_syrup.runtime.attention_region_phrase_evidence import ( + specific_attention_head_weights, +) +from simple_syrup.runtime.attention_region_self_completion import ( + ATTENTION_REGION_SELF_COMPLETION, + MAXIMUM_SELF_COMPLETION_ANCHORS, + SpatialSelfAttention, + _grid_anchor_indices, +) + + +def test_capture_pairs_spatial_self_attention_with_following_cross_call( + monkeypatch: Any, +) -> None: + """Pass same-layer self-attention only into derived concept evidence.""" + + session = _session(profile=AttentionCaptureProfile.EXHAUSTIVE) + observed_self_attention: list[SpatialSelfAttention | None] = [] + original = ATTENTION_AFFINITY_CALCULATOR.capture_spans + + def record_capture(*args: Any, **kwargs: Any) -> Any: + """Record staged self-attention before delegating to affinity capture.""" + + observed_self_attention.append(kwargs.get("self_attention")) + return original(*args, **kwargs) + + monkeypatch.setattr(ATTENTION_AFFINITY_CALCULATOR, "capture_spans", record_capture) + spatial = torch.tensor([[[1.0, 0.0], [0.9, 0.1], [0.8, 0.2], [0.0, 1.0]]]) + session.observe( + spatial, + spatial, + spatial, + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + session.observe( + spatial, + torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 1.0]]]), + torch.ones(1, 3, 2), + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + assert len(observed_self_attention) == 1 + assert isinstance(observed_self_attention[0], SpatialSelfAttention) + + +def test_anima_capture_does_not_apply_sdxl_self_attention_completion( + monkeypatch: Any, +) -> None: + """Keep Anima concept capture on its validated cross-attention evidence path.""" + + session = _session( + profile=AttentionCaptureProfile.EXHAUSTIVE, + sequence_length=512, + model_family=RegionalModelFamily.ANIMA, + source_aspect=0.75, + ) + observed_self_attention: list[SpatialSelfAttention | None] = [] + original = ATTENTION_AFFINITY_CALCULATOR.capture_spans + + def record_capture(*args: Any, **kwargs: Any) -> Any: + """Record completion input before delegating to affinity capture.""" + + observed_self_attention.append(kwargs.get("self_attention")) + return original(*args, **kwargs) + + monkeypatch.setattr(ATTENTION_AFFINITY_CALCULATOR, "capture_spans", record_capture) + spatial = torch.ones(1, 12, 2) + session.observe( + spatial, + spatial, + spatial, + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + session.observe( + spatial, + torch.ones(1, 512, 2), + torch.ones(1, 512, 2), + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + assert observed_self_attention == [None] + + +def test_anima_concept_evidence_preserves_contextualized_object_head() -> None: + """Let phrase modifiers refine Anima object evidence without erasing its extent.""" + + span = AttentionTokenSpan("pink hair", 1, (1, 2), (2,)) + session = _session( + profile=AttentionCaptureProfile.EXHAUSTIVE, + sequence_length=512, + model_family=RegionalModelFamily.ANIMA, + source_aspect=0.75, + catalog_spans=(span,), + ) + query = torch.tensor([[[2.0, 0.0], [0.0, 2.0], [1.5, 1.5]]]) + key = torch.zeros(1, 512, 2) + key[0, 1] = torch.tensor([1.0, 0.0]) + key[0, 2] = torch.tensor([0.0, 1.0]) + + session.observe( + query, + key, + torch.ones_like(key), + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + attention_map = session.maps_for("search")[0] + assert attention_map.concept_values is not None + assert attention_map.concept_values[1] > attention_map.concept_values[0] + assert attention_map.concept_values[1] > (attention_map.concept_values[2] * 0.75) + + +def test_anima_object_head_selection_rejects_diffuse_attention_heads() -> None: + """Prefer a spatially specific object head over broad interaction context.""" + + values = torch.tensor( + [ + [ + [0.1, 0.1, 0.9, 0.8], + [0.5, 0.5, 0.5, 0.5], + [0.4, 0.4, 0.4, 0.4], + [0.3, 0.3, 0.3, 0.3], + ] + ] + ) + + weights = specific_attention_head_weights(values) + + assert weights.shape == (1, 4) + assert weights[0, 0].item() == 1.0 + assert weights[0, 1:].count_nonzero().item() == 0 + + +def test_self_completion_drops_weak_grid_anchors() -> None: + """Keep spatial diversity without letting weak background cells steer completion.""" + + seed = torch.arange(9, dtype=torch.float32).reshape(1, 9) + + anchors = _grid_anchor_indices(seed, 3, 3) + + assert anchors.shape == (1, MAXIMUM_SELF_COMPLETION_ANCHORS) + assert set(anchors[0].tolist()) == {3, 4, 5, 6, 7, 8} + + +def test_self_completion_gates_related_recall_by_exact_object_grouping() -> None: + """Admit a related strand while rejecting an equally strong unrelated region.""" + + exact = torch.tensor([[1.0, 0.1, 0.05, 0.05]]) + related = torch.tensor([[1.0, 0.1, 0.9, 0.9]]) + query = torch.tensor([[[[1.0, 0.0], [0.0, 1.0], [1.0, 0.0], [0.0, 1.0]]]]) + self_attention = SpatialSelfAttention(query=query, key=query) + + completed = ATTENTION_REGION_SELF_COMPLETION.complete( + exact, + self_attention, + 2, + 2, + related_seed=related, + ) + + assert completed[0, 2] > completed[0, 3] + + +def test_sibling_consumers_share_one_concurrent_materialization( + monkeypatch: Any, +) -> None: + """Prevent parallel downstream nodes from consuming an emptied capture.""" + + session = _session(profile=AttentionCaptureProfile.EXHAUSTIVE) + session.observe( + torch.ones(1, 2, 2), + torch.ones(1, 3, 2), + torch.ones(1, 3, 2), + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + original = ATTENTION_AFFINITY_CALCULATOR.materialize + started = Event() + release = Event() + call_count = 0 + + def delayed_materialize(*args: Any, **kwargs: Any) -> Any: + """Hold the first materializer until a sibling is waiting.""" + + nonlocal call_count + call_count += 1 + started.set() + assert release.wait(timeout=2.0) + return original(*args, **kwargs) + + monkeypatch.setattr( + ATTENTION_AFFINITY_CALCULATOR, + "materialize", + delayed_materialize, + ) + with ThreadPoolExecutor(max_workers=2) as executor: + first = executor.submit(session.maps_for, "search") + assert started.wait(timeout=2.0) + second = executor.submit(session.maps_for, "search") + release.set() + results = (first.result(timeout=2.0), second.result(timeout=2.0)) + + assert call_count == 1 + assert all(len(result) == 1 for result in results) + assert results[0] is not results[1] + assert results[0][0].values.equal(results[1][0].values) + + +def test_capture_profile_subsamples_calls_without_changing_attention_output() -> None: + """Delegate every denoising call while retaining only fast-profile samples.""" + + session = _session(profile=AttentionCaptureProfile.FAST) + override = OptimizedAttentionCaptureOverride(session, None) + query = torch.ones(1, 2, 2) + key = torch.ones(1, 3, 2) + value = torch.arange(6, dtype=torch.float32).reshape(1, 3, 2) + + def original(*args: object, **kwargs: object) -> torch.Tensor: + """Return a sentinel output while accepting Comfy attention arguments.""" + + del args, kwargs + return torch.full((1, 2, 2), 7.0) + + options = tuple( + _options(cond_or_uncond=[0], block_index=index) for index in range(33) + ) + outputs = tuple( + override( + original, + query, + key, + value, + 1, + transformer_options=options[index], + ) + for index in range(33) + ) + + assert all(torch.equal(output, outputs[0]) for output in outputs) + assert outputs[0].eq(7.0).all().item() + assert len(session.maps_for("search")) == 2 + + +def test_fast_capture_rotates_sampled_layers_between_denoising_steps() -> None: + """Cover a different layer offset at each step without increasing stride cost.""" + + session = _session(profile=AttentionCaptureProfile.FAST) + query = torch.ones(1, 2, 2) + key = torch.ones(1, 3, 2) + value = torch.ones(1, 3, 2) + for block_index in range(33): + session.observe( + query, + key, + value, + 1, + _options(cond_or_uncond=[0], sigma=1.0, block_index=block_index), + skip_reshape=False, + ) + for block_index in range(2): + session.observe( + query, + key, + value, + 1, + _options(cond_or_uncond=[0], sigma=0.5, block_index=block_index), + skip_reshape=False, + ) + + assert len(session.maps_for("search")) == 3 + + +def test_fast_capture_subsamples_spatial_self_completion( + monkeypatch: Any, +) -> None: + """Retain sparse self-attention recall without paying for every fast sample.""" + + session = _session(profile=AttentionCaptureProfile.FAST) + observed: list[SpatialSelfAttention | None] = [] + original = ATTENTION_AFFINITY_CALCULATOR.capture_spans + + def capture_spans(*args: Any, **kwargs: Any) -> Any: + """Record completion inputs while preserving affinity behavior.""" + + observed.append(kwargs.get("self_attention")) + return original(*args, **kwargs) + + monkeypatch.setattr( + ATTENTION_AFFINITY_CALCULATOR, + "capture_spans", + capture_spans, + ) + for block_index in range(65): + options = _options(cond_or_uncond=[0], block_index=block_index) + session.observe( + torch.ones(1, 2, 2), + torch.ones(1, 2, 2), + torch.ones(1, 2, 2), + 1, + options, + skip_reshape=False, + ) + session.observe( + torch.ones(1, 2, 2), + torch.ones(1, 3, 2), + torch.ones(1, 3, 2), + 1, + options, + skip_reshape=False, + ) + + assert len(observed) == 3 + assert sum(isinstance(value, SpatialSelfAttention) for value in observed) == 1 + + +def test_coalesced_requests_share_one_unioned_native_affinity_pass( + monkeypatch: Any, +) -> None: + """Calculate shared prompt-token affinities once and route maps per request.""" + + controls = AttentionRegionControls( + 0.0, + 1.0, + 0.3, + 0.2, + 0.5, + 1, + AttentionCaptureProfile.EXHAUSTIVE, + ) + requests = ( + AttentionRegionRequest( + "all", + AttentionRegionRequestKind.ALL_PROMPT_SEGS, + (), + controls, + ), + AttentionRegionRequest( + "hair", + AttentionRegionRequestKind.CONCEPT_SEGS, + ("pink hair",), + controls, + ), + ) + plan = AttentionCapturePlan( + "sampler", + "sampler", + "model", + ("model", 0), + ("positive", 0), + requests, + "1girl, pink hair", + ("loader", 1), + ) + spans = ( + AttentionTokenSpan("1girl", 1, (1,)), + AttentionTokenSpan("pink hair", 1, (2,)), + ) + session = AttentionRegionCaptureSession( + plan=plan, + model_family=RegionalModelFamily.STANDARD_UNET, + token_catalog=AttentionTokenCatalog(4, spans, (0, 1, 2, 3)), + request_spans={"hair": (spans[1],), "all": spans}, + ) + original = ATTENTION_AFFINITY_CALCULATOR.capture_spans + calls: list[tuple[AttentionTokenSpan, ...]] = [] + + def record_capture(*args: Any, **kwargs: Any) -> Any: + """Record the unioned spans before delegating to real affinity math.""" + + captured_spans = kwargs.get("spans", args[3]) + calls.append(captured_spans) + return original(*args, **kwargs) + + monkeypatch.setattr(ATTENTION_AFFINITY_CALCULATOR, "capture_spans", record_capture) + session.observe( + torch.ones(1, 2, 2), + torch.ones(1, 4, 2), + torch.ones(1, 4, 2), + 1, + _options(cond_or_uncond=[0]), + skip_reshape=False, + ) + + assert calls == [spans] + assert tuple(value.label for value in session.maps_for("hair")) == ("pink hair",) + assert tuple(value.label for value in session.maps_for("all")) == ( + "1girl", + "pink hair", + ) + + +def _session( + profile: AttentionCaptureProfile = AttentionCaptureProfile.EXHAUSTIVE, + sequence_length: int = 3, + source_aspect: float | None = None, + model_family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, + token_indices: tuple[int, ...] = (1,), + catalog_spans: tuple[AttentionTokenSpan, ...] | None = None, +) -> AttentionRegionCaptureSession: + """Return one exact-native-query capture session.""" + + controls = AttentionRegionControls(0.0, 1.0, 0.3, 0.2, 0.5, 1, profile) + request = AttentionRegionRequest( + "search", + AttentionRegionRequestKind.CONCEPT_SEGS, + ("pink hair",), + controls, + ) + plan = AttentionCapturePlan( + "sampler", + "sampler", + "model", + ("model", 0), + ("positive", 0), + (request,), + "pink hair", + ("loader", 1), + source_aspect, + ) + default_span = AttentionTokenSpan("pink hair", 1, token_indices) + catalog = AttentionTokenCatalog( + sequence_length, + catalog_spans or (default_span,), + tuple(range(sequence_length)), + ) + return AttentionRegionCaptureSession( + plan=plan, + model_family=model_family, + token_catalog=catalog, + request_spans={"search": (catalog.spans[0],)}, + ) + + +def _options( + *, + cond_or_uncond: list[int], + sigma: float = 1.0, + block_index: int = 0, +) -> dict[str, object]: + """Return exact sampler metadata at the beginning of denoising.""" + + return { + "sample_sigmas": torch.tensor([1.0, 0.5, 0.0]), + "sigmas": torch.tensor([sigma]), + "cond_or_uncond": cond_or_uncond, + "block": ("middle", 0), + "block_index": block_index, + } + + +def _patcher() -> Any: + """Create a real CPU ModelPatcher with isolated transformer options.""" + + from comfy.model_patcher import ModelPatcher + + base_model = torch.nn.Module() + base_model.diffusion_model = torch.nn.Linear(1, 1) + device = torch.device("cpu") + return ModelPatcher(base_model, load_device=device, offload_device=device) diff --git a/tests/regional_generation/attention_regions/test_attention_region_components.py b/tests/regional_generation/attention_regions/test_attention_region_components.py new file mode 100644 index 0000000..00e8a37 --- /dev/null +++ b/tests/regional_generation/attention_regions/test_attention_region_components.py @@ -0,0 +1,508 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Test attention-native temporal shaping and soft SEGS construction.""" + +from __future__ import annotations + +import torch + +from simple_syrup.domain.attention_region_capture import ( + AttentionCaptureProfile, + AttentionEvidenceMode, + AttentionRegionControls, +) +from simple_syrup.domain.attention_region_maps import CapturedAttentionMap +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily +from simple_syrup.services.attention_region_matte import ATTENTION_MATTE_SERVICE +from simple_syrup.services.attention_region_rendering import ( + ATTENTION_REGION_RENDERING_SERVICE, +) + + +def test_empty_maps_return_correctly_sized_no_op_outputs() -> None: + """Return empty SEGS and a zero mask for unsupported graph/model paths.""" + + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(), + controls=_controls(), + height=5, + width=7, + ) + + assert segs == ((5, 7), ()) + assert mask.shape == (1, 5, 7) + assert mask.count_nonzero().item() == 0 + + +def test_disconnected_attention_islands_become_separate_instances() -> None: + """Expose retained disconnected objects as independently controllable SEGS.""" + + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("outdoors", [1.0, 0.0, 1.0], 0.5),), + controls=_controls(strength=0.5, consensus=0.0), + height=1, + width=3, + ) + + assert len(segs[1]) == 2 + assert tuple(segment.label for segment in segs[1]) == ("outdoors", "outdoors") + assert mask.count_nonzero().item() == 2 + + +def test_concept_isolation_rejects_weak_disconnected_context() -> None: + """Drop a weak contextual island while retaining raw inspection evidence.""" + + maps = (_map("bangs", [1.0, 0.9, 0.0, 0.3, 0.3], 0.5),) + concept_segs, _concept_mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls(strength=0.15, consensus=0.0), + height=1, + width=5, + ) + raw_segs, _raw_mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.0, + evidence_mode=AttentionEvidenceMode.RAW, + ), + height=1, + width=5, + ) + + assert len(concept_segs[1]) == 1 + assert concept_segs[1][0].bbox == (0, 0, 2, 1) + assert len(raw_segs[1]) == 2 + + +def test_concept_isolation_retains_multiple_confident_regions() -> None: + """Keep plural concept instances when each has substantial evidence.""" + + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("cuffs", [1.0, 0.9, 0.0, 0.7, 0.65], 0.5),), + controls=_controls(strength=0.15, consensus=0.0), + height=1, + width=5, + ) + + assert len(segs[1]) == 2 + assert mask.count_nonzero().item() == 4 + + +def test_concept_isolation_retains_sparse_instances_by_peak_evidence() -> None: + """Keep small repeated instances without rewarding a larger region for area.""" + + values = [0.6] * 16 + [0.0, 1.0, 0.0, 0.7] + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("falling petals", values, 0.5),), + controls=_controls(strength=0.15, consensus=0.0), + height=1, + width=len(values), + ) + + assert len(segs[1]) == 3 + assert mask[0, 0, 17].item() > 0.0 + assert mask[0, 0, 19].item() > 0.0 + + +def test_instance_recall_can_restrict_sparse_results_to_the_strongest_peak() -> None: + """Let users remove weaker disconnected instances without an area heuristic.""" + + values = [0.6] * 16 + [0.0, 1.0, 0.0, 0.7] + segs, _mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("falling petals", values, 0.5),), + controls=_controls( + strength=0.15, + consensus=0.0, + instance_recall=0.2, + ), + height=1, + width=len(values), + ) + + assert len(segs[1]) == 1 + assert segs[1][0].bbox == (17, 0, 18, 1) + + +def test_concept_isolation_does_not_prefer_a_tiny_sharp_island() -> None: + """Keep broad supported evidence when a disconnected pixel peaks higher.""" + + values = [0.4] * 25 + [0.0, 1.0] + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("hair", values, 0.5),), + controls=_controls(strength=0.15, consensus=0.0), + height=1, + width=len(values), + ) + + assert any(segment.bbox == (0, 0, 25, 1) for segment in segs[1]) + assert mask[0, 0, :25].count_nonzero().item() == 25 + + +def test_concept_isolation_rejects_pockmarks_around_a_compact_dominant_region() -> None: + """Keep one compact body when much smaller disconnected peaks surround it.""" + + values = torch.zeros(11, 11) + values[3:8, 3:8] = 0.6 + values[0:2, 0:2] = 0.8 + values[0, 10] = 1.0 + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("torso", values.flatten().tolist(), 0.5),), + controls=_controls(strength=0.15, consensus=0.0, feather=0), + height=11, + width=11, + ) + + assert len(segs[1]) == 1 + assert segs[1][0].bbox == (3, 3, 8, 8) + assert mask.count_nonzero().item() == 25 + + +def test_full_instance_recall_preserves_pockmarks_around_a_compact_region() -> None: + """Honor an explicit request to retain every supported disconnected instance.""" + + values = torch.zeros(11, 11) + values[3:8, 3:8] = 0.6 + values[0:2, 0:2] = 0.8 + values[0, 10] = 1.0 + segs, _mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("torso", values.flatten().tolist(), 0.5),), + controls=_controls( + strength=0.15, + consensus=0.0, + feather=0, + instance_recall=1.0, + ), + height=11, + width=11, + ) + + assert len(segs[1]) == 3 + + +def test_split_sensitivity_preserves_the_complete_concept_union() -> None: + """Change instance separation without deleting moderate concept support.""" + + maps = (_map("cat", [1.0, 0.4, 0.4, 0.4, 1.0], 0.5),) + _unsplit_segs, unsplit = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls(strength=0.3, consensus=0.0, split=0.0), + height=1, + width=5, + ) + split_segs, split = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls(strength=0.3, consensus=0.0, split=1.0), + height=1, + width=5, + ) + + assert torch.equal(split, unsplit) + assert split.count_nonzero().item() == 5 + assert len(split_segs[1]) == 2 + + +def test_cohesive_support_retains_a_moderate_body_around_a_strong_core() -> None: + """Keep the complete attended silhouette instead of its sparse core pixels.""" + + values = [ + 0.0, + 0.2, + 0.2, + 0.2, + 0.0, + 0.2, + 0.4, + 0.6, + 0.4, + 0.2, + 0.2, + 0.6, + 1.0, + 0.6, + 0.2, + 0.2, + 0.4, + 0.6, + 0.4, + 0.2, + 0.0, + 0.2, + 0.2, + 0.2, + 0.0, + ] + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("hair", values, 0.5),), + controls=_controls(strength=0.15, consensus=0.0, split=0.75), + height=5, + width=5, + ) + + assert mask.count_nonzero().item() == 21 + assert mask[0, 2, 2].item() == 1.0 + assert mask[0, 0, 1].item() > 0.0 + + +def test_minimum_region_size_removes_each_pockmark_independently() -> None: + """Discard a tiny island without rejecting or merging the valid component.""" + + values = [1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0] + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("hair", values, 0.5),), + controls=_controls(strength=0.5, consensus=0.0, minimum_size=2), + height=3, + width=3, + ) + + assert len(segs[1]) == 1 + assert segs[1][0].bbox == (0, 0, 2, 1) + assert mask.count_nonzero().item() == 2 + + +def test_keep_only_and_combine_apply_per_concept_without_changing_union() -> None: + """Retain top instances and optionally package their union as one SEG.""" + + maps = (_map("cat", [1.0, 1.0, 0.0, 0.8, 0.8], 0.5),) + separate, separate_mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.5, + consensus=0.0, + keep_only=2, + combine=False, + ), + height=1, + width=5, + ) + combined, combined_mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.5, + consensus=0.0, + keep_only=2, + combine=True, + ), + height=1, + width=5, + ) + + assert len(separate[1]) == 2 + assert len(combined[1]) == 1 + assert torch.equal(separate_mask, combined_mask) + + +def test_keep_largest_groups_a_nearby_detached_concept_fragment() -> None: + """Treat qualifying nearby support as one instance before top-N ranking.""" + + values = torch.zeros((9, 9), dtype=torch.float32) + values[1, 1:8] = 1.0 + values[7, 1:8] = 1.0 + values[1:8, 1] = 1.0 + values[4:6, 4:6] = 0.8 + + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("cat", values.flatten().tolist(), 0.5),), + controls=_controls( + strength=0.5, + consensus=0.0, + keep_only=1, + feather=0, + ), + height=9, + width=9, + ) + + assert len(segs[1]) == 1 + assert mask[0, 4:6, 4:6].count_nonzero().item() == 4 + + +def test_matte_solidity_flattens_interior_and_edge_feather_softens_boundary() -> None: + """Preserve raw alpha at zero and make a solid feathered matte at one.""" + + maps = (_map("dress", [0.0, 0.6, 1.0, 0.7, 0.0], 0.5),) + _raw_segs, raw = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls(strength=0.2, consensus=0.0, solidity=0.0), + height=1, + width=5, + ) + _solid_segs, solid = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.2, + consensus=0.0, + solidity=1.0, + feather=1, + ), + height=1, + width=5, + ) + + assert raw[0, 0, 1].item() != raw[0, 0, 2].item() + assert solid[0, 0, 2].item() == 1.0 + assert 0.0 < solid[0, 0, 0].item() < 1.0 + + +def test_full_matte_solidity_fills_only_enclosed_holes() -> None: + """Fill an interior attention gap without filling exterior-connected space.""" + + ring = [1.0, 1.0, 1.0, 1.0, 0.0, 1.0, 1.0, 1.0, 1.0] + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("hair", ring, 0.5),), + controls=_controls( + strength=0.5, + consensus=0.0, + solidity=1.0, + feather=0, + ), + height=3, + width=3, + ) + + assert torch.equal(mask, torch.ones_like(mask)) + + +def test_full_matte_solidity_preserves_a_narrow_exterior_connected_channel() -> None: + """Flatten alpha without inventing support inside an exterior channel.""" + + support = torch.zeros(7, 7, dtype=torch.bool) + support[1:6, 1:6] = True + support[1:4, 3] = False + + matte = ATTENTION_MATTE_SERVICE.shape( + alpha=support.float(), + support=support, + solidity=1.0, + edge_feather=0, + ) + + assert matte[3, 3].item() == 0.0 + assert matte[0].count_nonzero().item() == 0 + assert matte[:, 0].count_nonzero().item() == 0 + + +def test_full_matte_solidity_preserves_a_winding_exterior_channel() -> None: + """Keep exterior-connected exclusions regardless of their shape or width.""" + + support = torch.ones(100, 100, dtype=torch.bool) + support[55:, 25:75] = False + support[35:55, 49:52] = False + support[42:45, 42:52] = False + + matte = ATTENTION_MATTE_SERVICE.shape( + alpha=support.float(), + support=support, + solidity=1.0, + edge_feather=0, + ) + + assert matte[38, 50].item() == 0.0 + assert matte[43, 44].item() == 0.0 + assert matte[80, 50].item() == 0.0 + + +def test_full_matte_solidity_never_erases_accepted_thin_support() -> None: + """Keep every accepted pixel when topology cleanup adds cohesive support.""" + + support = torch.zeros(100, 100, dtype=torch.bool) + support[10:90, 50] = True + + matte = ATTENTION_MATTE_SERVICE.shape( + alpha=support.float(), + support=support, + solidity=1.0, + edge_feather=0, + ) + + assert torch.all(matte[support] == 1.0) + + +def test_highest_confidence_keeps_stronger_component() -> None: + """Rank components from pre-normalized evidence rather than normalized maxima.""" + + segs, _mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=(_map("cat", [1.0, 0.0, 0.55], 0.5),), + controls=AttentionRegionControls( + 0.0, + 1.0, + 0.4, + 0.0, + 0.0, + 1, + AttentionCaptureProfile.BALANCED, + keep_only=1, + keep_by="highest confidence", + ), + height=1, + width=3, + ) + + assert len(segs[1]) == 1 + assert segs[1][0].bbox == (0, 0, 1, 1) + + +def _map( + label: str, + values: list[float], + progress: float, + *, + baseline: float = 0.0, + layer: str = "layer", + family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, + concept_values: list[float] | None = None, +) -> CapturedAttentionMap: + """Create one compact square attention observation.""" + + return CapturedAttentionMap( + label, + torch.tensor(values, dtype=torch.float16), + progress, + layer, + uniform_probability=baseline, + model_family=family, + concept_values=( + torch.tensor(concept_values, dtype=torch.float16) + if concept_values is not None + else None + ), + ) + + +def _controls( + *, + start: float = 0.0, + end: float = 1.0, + strength: float = 0.35, + consensus: float = 0.25, + minimum_size: int = 1, + keep_only: int = 0, + combine: bool = False, + solidity: float = 0.0, + feather: int = 8, + split: float = 0.0, + instance_recall: float = 0.65, + geometry_recall: float = 0.85, + evidence_mode: AttentionEvidenceMode = AttentionEvidenceMode.CONCEPT, +) -> AttentionRegionControls: + """Return representative balanced rendering controls.""" + + return AttentionRegionControls( + start, + end, + strength, + consensus, + split, + minimum_size, + AttentionCaptureProfile.BALANCED, + instance_recall=instance_recall, + geometry_recall=geometry_recall, + keep_only=keep_only, + combine_segs=combine, + matte_solidity=solidity, + edge_feather=feather, + evidence_mode=evidence_mode, + ) diff --git a/tests/regional_generation/attention_regions/test_attention_region_evidence.py b/tests/regional_generation/attention_regions/test_attention_region_evidence.py new file mode 100644 index 0000000..eb251b7 --- /dev/null +++ b/tests/regional_generation/attention_regions/test_attention_region_evidence.py @@ -0,0 +1,332 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Test attention-native temporal shaping and soft SEGS construction.""" + +from __future__ import annotations + +import torch + +from simple_syrup.domain.attention_region_capture import ( + AttentionCaptureProfile, + AttentionEvidenceMode, + AttentionRegionControls, +) +from simple_syrup.domain.attention_region_maps import CapturedAttentionMap +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily +from simple_syrup.services.attention_region_rendering import ( + ATTENTION_REGION_RENDERING_SERVICE, +) + + +def test_geometry_recall_controls_faint_connected_exact_token_extent() -> None: + """Let users trade faint attached geometry for a tighter semantic core.""" + + maps = tuple( + _map( + "cat tail", + [ + 1.0, + 1.0, + 0.6, + 0.6, + 0.4, + 0.4, + 0.2, + 0.2, + 0.0, + 0.0, + 0.0, + 0.0, + 0.0, + 0.0, + 0.0, + 0.0, + ], + progress, + family=RegionalModelFamily.ANIMA, + concept_values=[1.0, 1.0] + [0.0] * 14, + ) + for progress in (0.2, 0.7) + ) + + _recalled_segs, recalled = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=1.0, + ), + height=1, + width=16, + ) + _tight_segs, tight = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=0.0, + ), + height=1, + width=16, + ) + + assert recalled.count_nonzero().item() == 8 + assert tight.count_nonzero().item() == 6 + + +def test_anima_concept_mode_rejects_below_baseline_late_residue() -> None: + """Do not normalize negligible late-step Anima residue into full support.""" + + maps = tuple( + _map( + "cuffs", + [0.01, 0.01, 0.01, 0.01], + progress, + family=RegionalModelFamily.ANIMA, + concept_values=[0.0010, 0.0012, 0.0011, 0.0010], + baseline=0.01, + ) + for progress in (0.75, 0.9) + ) + + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=2, + width=2, + ) + + assert segs[1] == () + assert mask.count_nonzero().item() == 0 + + +def test_anima_concept_mode_removes_a_broad_contextual_field() -> None: + """Keep local lift without treating a broadly elevated phrase as full-frame.""" + + maps = tuple( + _map( + "swept bangs", + [0.1, 0.1, 0.1, 0.1], + progress, + family=RegionalModelFamily.ANIMA, + concept_values=[0.11, 0.11, 0.11, 0.14], + baseline=0.1, + ) + for progress in (0.2, 0.6) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=2, + width=2, + ) + + assert mask.count_nonzero().item() == 1 + assert mask[0, 1, 1].item() == 1.0 + + +def test_anima_concept_mode_prefers_resolved_mid_pass_evidence() -> None: + """Prevent an endpoint observation from tying resolved middle evidence.""" + + maps = ( + _map( + "boots", + [0.1, 0.1, 0.1, 0.1], + 0.0, + family=RegionalModelFamily.ANIMA, + concept_values=[0.9, 0.1, 0.1, 0.1], + baseline=0.1, + layer="shared", + ), + _map( + "boots", + [0.1, 0.1, 0.1, 0.1], + 0.5, + family=RegionalModelFamily.ANIMA, + concept_values=[0.1, 0.1, 0.1, 0.9], + baseline=0.1, + layer="shared", + ), + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.4, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=2, + width=2, + ) + + assert mask.count_nonzero().item() == 1 + assert mask[0, 1, 1].item() == 1.0 + + +def test_concept_isolation_prefers_repeated_support_over_transient_peak() -> None: + """Suppress one strong flash when stable observations localize elsewhere.""" + + maps = ( + _map( + "hair", + [0.25, 0.25, 0.25, 0.95], + 0.1, + baseline=0.25, + layer="early", + ), + _map( + "hair", + [0.25, 0.78, 0.25, 0.25], + 0.4, + baseline=0.25, + layer="middle-a", + ), + _map( + "hair", + [0.25, 0.82, 0.25, 0.25], + 0.6, + baseline=0.25, + layer="middle-b", + ), + _map( + "hair", + [0.25, 0.76, 0.25, 0.25], + 0.8, + baseline=0.25, + layer="late", + ), + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.2, + consensus=0.4, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=2, + width=2, + ) + + assert mask.count_nonzero().item() == 1 + assert mask[0, 0, 1].item() == 1.0 + + +def test_concept_isolation_preserves_repeated_fine_layer_structure() -> None: + """Keep layer-local fine detail without admitting a one-step transient.""" + + maps = tuple( + _map( + "hair", + [0.25, 0.85, 0.25, 0.25], + progress, + baseline=0.25, + layer="coarse", + ) + for progress in (0.2, 0.4, 0.6, 0.8) + ) + ( + _map( + "hair", + [0.25, 0.85, 0.52, 0.48], + 0.4, + baseline=0.25, + layer="fine", + ), + _map( + "hair", + [0.25, 0.85, 0.55, 0.25], + 0.7, + baseline=0.25, + layer="fine", + ), + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.5, + feather=0, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=2, + width=2, + ) + + assert mask[0, 1, 0].item() > 0.0 + assert mask[0, 1, 1].item() == 0.0 + + +def _map( + label: str, + values: list[float], + progress: float, + *, + baseline: float = 0.0, + layer: str = "layer", + family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, + concept_values: list[float] | None = None, +) -> CapturedAttentionMap: + """Create one compact square attention observation.""" + + return CapturedAttentionMap( + label, + torch.tensor(values, dtype=torch.float16), + progress, + layer, + uniform_probability=baseline, + model_family=family, + concept_values=( + torch.tensor(concept_values, dtype=torch.float16) + if concept_values is not None + else None + ), + ) + + +def _controls( + *, + start: float = 0.0, + end: float = 1.0, + strength: float = 0.35, + consensus: float = 0.25, + minimum_size: int = 1, + keep_only: int = 0, + combine: bool = False, + solidity: float = 0.0, + feather: int = 8, + split: float = 0.0, + instance_recall: float = 0.65, + geometry_recall: float = 0.85, + evidence_mode: AttentionEvidenceMode = AttentionEvidenceMode.CONCEPT, +) -> AttentionRegionControls: + """Return representative balanced rendering controls.""" + + return AttentionRegionControls( + start, + end, + strength, + consensus, + split, + minimum_size, + AttentionCaptureProfile.BALANCED, + instance_recall=instance_recall, + geometry_recall=geometry_recall, + keep_only=keep_only, + combine_segs=combine, + matte_solidity=solidity, + edge_feather=feather, + evidence_mode=evidence_mode, + ) diff --git a/tests/regional_generation/attention_regions/test_attention_region_geometry.py b/tests/regional_generation/attention_regions/test_attention_region_geometry.py new file mode 100644 index 0000000..427b916 --- /dev/null +++ b/tests/regional_generation/attention_regions/test_attention_region_geometry.py @@ -0,0 +1,500 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Test attention-native temporal shaping and soft SEGS construction.""" + +from __future__ import annotations + +import torch + +from simple_syrup.domain.attention_region_capture import ( + AttentionCaptureProfile, + AttentionEvidenceMode, + AttentionRegionControls, +) +from simple_syrup.domain.attention_region_maps import CapturedAttentionMap +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily +from simple_syrup.services.attention_region_rendering import ( + ATTENTION_REGION_RENDERING_SERVICE, +) + + +def test_anima_concept_mode_uses_validated_probability_aggregation() -> None: + """Keep Anima on its spatially faithful cross-attention evidence policy.""" + + maps = ( + _map( + "cat", + [0.1, 0.8, 0.2, 0.7], + 0.2, + family=RegionalModelFamily.ANIMA, + concept_values=[0.8, 0.1, 0.7, 0.2], + baseline=0.1, + ), + _map( + "cat", + [0.1, 0.9, 0.1, 0.2], + 0.7, + family=RegionalModelFamily.ANIMA, + concept_values=[0.9, 0.1, 0.8, 0.1], + baseline=0.1, + ), + ) + concept_controls = _controls( + strength=0.2, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ) + raw_controls = _controls( + strength=0.2, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.RAW, + ) + + _concept_segs, concept_mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=concept_controls, + height=2, + width=2, + ) + _raw_segs, raw_mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=raw_controls, + height=2, + width=2, + ) + + assert not torch.equal(concept_mask, raw_mask) + assert concept_mask[0, 0, 0] > concept_mask[0, 0, 1] + assert raw_mask[0, 0, 1] > raw_mask[0, 0, 0] + + +def test_anima_concept_mode_recovers_connected_exact_token_geometry() -> None: + """Keep a raw-attention appendage attached to the semantic concept core.""" + + raw = [0.0] * 5 + [1.0, 1.0, 0.8, 0.7, 0.0] + [0.0] * 5 + concept = [0.0] * 5 + [1.0, 1.0, 0.0, 0.0, 0.0] + [0.0] * 5 + maps = tuple( + _map( + "cat", + raw, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=concept, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.5, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=3, + width=5, + ) + + assert mask[0, 1, :4].count_nonzero().item() == 4 + assert mask.count_nonzero().item() == 4 + + +def test_anima_concept_mode_rejects_detached_exact_token_noise() -> None: + """Exclude raw-attention components that do not touch the semantic core.""" + + maps = tuple( + _map( + "cat", + [1.0, 1.0, 0.0, 0.0, 0.9], + progress, + family=RegionalModelFamily.ANIMA, + concept_values=[1.0, 1.0, 0.0, 0.0, 0.0], + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.5, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=1, + width=5, + ) + + assert mask[0, 0, :2].count_nonzero().item() == 2 + assert mask[0, 0, 4].item() == 0.0 + + +def test_anima_concept_mode_preserves_complete_expansive_anchored_geometry() -> None: + """Keep the complete attached object while dropping its weak global bridge.""" + + maps = tuple( + _map( + "mage staff", + [1.0, 1.0, 0.8, 0.8, 0.8, 0.2, 0.2, 0.2, 0.2, 0.2], + progress, + family=RegionalModelFamily.ANIMA, + concept_values=[1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0], + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=1, + width=10, + ) + + assert mask[0, 0, :5].count_nonzero().item() == 5 + assert mask[0, 0, 5:].count_nonzero().item() == 0 + + +def test_anima_concept_mode_favors_recall_when_expansion_is_ambiguous() -> None: + """Keep attached geometry when no tighter support extends beyond the core.""" + + maps = tuple( + _map( + "close subject", + [1.0, 1.0, 0.8, 0.8, 0.8], + progress, + family=RegionalModelFamily.ANIMA, + concept_values=[1.0, 1.0, 0.0, 0.0, 0.0], + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=1, + width=5, + ) + + assert mask.count_nonzero().item() == 5 + + +def test_anima_concept_mode_rejects_weak_broad_field_around_compact_peak() -> None: + """Keep a compact semantic peak without absorbing its weak connected field.""" + + raw = [0.0] * 100 + for row in range(3, 7): + for column in range(10): + raw[row * 10 + column] = 0.2 + raw[44] = 1.0 + raw[45] = 1.0 + concept = [0.0] * 100 + concept[44] = 1.0 + concept[45] = 1.0 + maps = tuple( + _map( + "compact feature", + raw, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=concept, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=0.85, + ), + height=10, + width=10, + ) + + assert mask.count_nonzero().item() == 2 + + +def test_anima_concept_mode_rejects_global_reference_geometry_for_local_core() -> None: + """Do not expand localized evidence through a near-global exact-token field.""" + + raw = [0.3] * 100 + concept = [0.0] * 100 + for row in range(3, 7): + for column in range(10): + raw[row * 10 + column] = 0.4 + for column in range(3, 8): + concept[4 * 10 + column] = 1.0 + maps = tuple( + _map( + "localized feature", + raw, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=concept, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=0.85, + ), + height=10, + width=10, + ) + + assert mask.count_nonzero().item() == 5 + + +def test_anima_concept_mode_prefers_concentrated_geometry_for_moderate_core() -> None: + """Use stricter geometry consensus when a moderate core anchors a broad field.""" + + raw = [0.0] * 100 + for row in range(3, 7): + for column in range(10): + raw[row * 10 + column] = 0.2 + concept = [0.0] * 100 + for column in range(10): + raw[4 * 10 + column] = 1.0 + concept[4 * 10 + column] = 1.0 + maps = tuple( + _map( + "moderate compact feature", + raw, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=concept, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=0.85, + ), + height=10, + width=10, + ) + + assert mask.count_nonzero().item() == 10 + + +def test_anima_concept_mode_tightens_compact_geometry_below_frame_threshold() -> None: + """Use the stable core when a compact weak field occupies under 15% of a frame.""" + + raw = [0.0] * 400 + for row in range(6, 13): + for column in range(6, 13): + raw[row * 20 + column] = 0.2 + concept = [0.0] * 400 + for row in range(8, 11): + for column in range(8, 11): + raw[row * 20 + column] = 1.0 + concept[row * 20 + column] = 1.0 + maps = tuple( + _map( + "compact sub-frame feature", + raw, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=concept, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=0.85, + ), + height=20, + width=20, + ) + + assert mask.count_nonzero().item() == 9 + + +def test_full_geometry_recall_preserves_weak_broad_field_around_compact_peak() -> None: + """Let an explicit maximum-recall choice bypass adaptive compactness.""" + + raw = [0.0] * 100 + for row in range(3, 7): + for column in range(10): + raw[row * 10 + column] = 0.2 + raw[44] = 1.0 + raw[45] = 1.0 + concept = [0.0] * 100 + concept[44] = 1.0 + concept[45] = 1.0 + maps = tuple( + _map( + "compact feature", + raw, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=concept, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=1.0, + ), + height=10, + width=10, + ) + + assert mask.count_nonzero().item() == 40 + + +def test_anima_concept_mode_tightens_peak_dominated_semantic_field() -> None: + """Tighten broad support when strength belongs mainly to a small core.""" + + values = [0.05] * 100 + for row in range(3, 7): + for column in range(10): + values[row * 10 + column] = 0.3 + values[44] = 1.0 + values[45] = 1.0 + maps = tuple( + _map( + "compact semantic feature", + values, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=values, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=0.85, + ), + height=10, + width=10, + ) + + assert mask.count_nonzero().item() == 2 + + +def test_anima_concept_mode_preserves_broad_high_strength_semantic_region() -> None: + """Keep a genuinely broad concept whose support remains strong across its area.""" + + values = [0.05] * 100 + for row in range(2, 8): + for column in range(10): + values[row * 10 + column] = 0.8 + values[44] = 1.0 + values[45] = 1.0 + maps = tuple( + _map( + "broad semantic region", + values, + progress, + family=RegionalModelFamily.ANIMA, + concept_values=values, + ) + for progress in (0.2, 0.7) + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.25, + geometry_recall=0.85, + ), + height=10, + width=10, + ) + + assert mask.count_nonzero().item() == 60 + + +def _map( + label: str, + values: list[float], + progress: float, + *, + baseline: float = 0.0, + layer: str = "layer", + family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, + concept_values: list[float] | None = None, +) -> CapturedAttentionMap: + """Create one compact square attention observation.""" + + return CapturedAttentionMap( + label, + torch.tensor(values, dtype=torch.float16), + progress, + layer, + uniform_probability=baseline, + model_family=family, + concept_values=( + torch.tensor(concept_values, dtype=torch.float16) + if concept_values is not None + else None + ), + ) + + +def _controls( + *, + start: float = 0.0, + end: float = 1.0, + strength: float = 0.35, + consensus: float = 0.25, + minimum_size: int = 1, + keep_only: int = 0, + combine: bool = False, + solidity: float = 0.0, + feather: int = 8, + split: float = 0.0, + instance_recall: float = 0.65, + geometry_recall: float = 0.85, + evidence_mode: AttentionEvidenceMode = AttentionEvidenceMode.CONCEPT, +) -> AttentionRegionControls: + """Return representative balanced rendering controls.""" + + return AttentionRegionControls( + start, + end, + strength, + consensus, + split, + minimum_size, + AttentionCaptureProfile.BALANCED, + instance_recall=instance_recall, + geometry_recall=geometry_recall, + keep_only=keep_only, + combine_segs=combine, + matte_solidity=solidity, + edge_feather=feather, + evidence_mode=evidence_mode, + ) diff --git a/tests/test_attention_region_graph.py b/tests/regional_generation/attention_regions/test_attention_region_graph.py similarity index 100% rename from tests/test_attention_region_graph.py rename to tests/regional_generation/attention_regions/test_attention_region_graph.py diff --git a/tests/test_attention_region_logits.py b/tests/regional_generation/attention_regions/test_attention_region_logits.py similarity index 100% rename from tests/test_attention_region_logits.py rename to tests/regional_generation/attention_regions/test_attention_region_logits.py diff --git a/tests/test_attention_region_prompt_handler.py b/tests/regional_generation/attention_regions/test_attention_region_prompt_handler.py similarity index 100% rename from tests/test_attention_region_prompt_handler.py rename to tests/regional_generation/attention_regions/test_attention_region_prompt_handler.py diff --git a/tests/regional_generation/attention_regions/test_attention_region_rendering.py b/tests/regional_generation/attention_regions/test_attention_region_rendering.py new file mode 100644 index 0000000..f19b6a7 --- /dev/null +++ b/tests/regional_generation/attention_regions/test_attention_region_rendering.py @@ -0,0 +1,207 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Test attention-native temporal shaping and soft SEGS construction.""" + +from __future__ import annotations + +import torch + +from simple_syrup.domain.attention_region_capture import ( + AttentionCaptureProfile, + AttentionEvidenceMode, + AttentionRegionControls, +) +from simple_syrup.domain.attention_region_maps import CapturedAttentionMap +from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily +from simple_syrup.services.attention_region_rendering import ( + ATTENTION_REGION_RENDERING_SERVICE, +) + + +def test_strength_and_consensus_shape_soft_regions_before_component_packaging() -> None: + """Reject transient weak pixels while preserving feathered accepted values.""" + + maps = ( + _map("hair", [0.1, 0.8, 0.2, 0.7], 0.2), + _map("hair", [0.1, 0.9, 0.1, 0.2], 0.5), + _map("hair", [0.1, 0.7, 0.1, 0.1], 0.8), + ) + image = torch.zeros(1, 2, 2, 3) + + segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls(strength=0.5, consensus=0.66), + height=2, + width=2, + image=image, + ) + + assert len(segs[1]) == 1 + assert segs[1][0].label == "hair" + assert isinstance(segs[1][0].cropped_mask, torch.Tensor) + assert segs[1][0].cropped_mask.shape == (1, 1) + assert mask[0, 0, 1].item() == 1.0 + assert mask.count_nonzero().item() == 1 + + +def test_temporal_window_changes_region_without_morphology() -> None: + """Select early versus late attention evidence through normalized progress.""" + + maps = ( + _map("composition", [1.0, 1.0, 0.0, 0.0], 0.1), + _map("composition", [0.0, 0.0, 1.0, 1.0], 0.9), + ) + + _early_segs, early = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls(start=0.0, end=0.4, strength=0.2, consensus=0.0), + height=2, + width=2, + ) + _late_segs, late = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls(start=0.6, end=1.0, strength=0.2, consensus=0.0), + height=2, + width=2, + ) + + assert early[0, 0].sum().item() > early[0, 1].sum().item() + assert late[0, 1].sum().item() > late[0, 0].sum().item() + + +def test_raw_attention_preserves_diffuse_model_evidence() -> None: + """Keep low-amplitude positive attention visible in inspection mode.""" + + maps = ( + _map( + "outfit", + [0.26, 0.25, 0.27, 0.25], + 0.1, + baseline=0.25, + ), + _map( + "outfit", + [0.25, 0.80, 0.25, 0.25], + 0.6, + baseline=0.25, + ), + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.0, + consensus=0.0, + evidence_mode=AttentionEvidenceMode.RAW, + ), + height=2, + width=2, + ) + + assert mask.count_nonzero().item() == 4 + assert mask[0, 0, 1].item() == 1.0 + assert mask[0, 0, 0].item() > 0.0 + + +def test_concept_isolation_removes_uniform_attention_baseline() -> None: + """Do not promote near-uniform attention into concept support.""" + + maps = ( + _map( + "outfit", + [0.26, 0.25, 0.27, 0.25], + 0.1, + baseline=0.25, + ), + _map( + "outfit", + [0.25, 0.80, 0.25, 0.25], + 0.5, + baseline=0.25, + ), + _map( + "outfit", + [0.25, 0.75, 0.25, 0.25], + 0.7, + baseline=0.25, + ), + ) + + _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( + maps=maps, + controls=_controls( + strength=0.15, + consensus=0.2, + evidence_mode=AttentionEvidenceMode.CONCEPT, + ), + height=2, + width=2, + ) + + assert mask.count_nonzero().item() == 1 + assert mask[0, 0, 1].item() == 1.0 + + +def _map( + label: str, + values: list[float], + progress: float, + *, + baseline: float = 0.0, + layer: str = "layer", + family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, + concept_values: list[float] | None = None, +) -> CapturedAttentionMap: + """Create one compact square attention observation.""" + + return CapturedAttentionMap( + label, + torch.tensor(values, dtype=torch.float16), + progress, + layer, + uniform_probability=baseline, + model_family=family, + concept_values=( + torch.tensor(concept_values, dtype=torch.float16) + if concept_values is not None + else None + ), + ) + + +def _controls( + *, + start: float = 0.0, + end: float = 1.0, + strength: float = 0.35, + consensus: float = 0.25, + minimum_size: int = 1, + keep_only: int = 0, + combine: bool = False, + solidity: float = 0.0, + feather: int = 8, + split: float = 0.0, + instance_recall: float = 0.65, + geometry_recall: float = 0.85, + evidence_mode: AttentionEvidenceMode = AttentionEvidenceMode.CONCEPT, +) -> AttentionRegionControls: + """Return representative balanced rendering controls.""" + + return AttentionRegionControls( + start, + end, + strength, + consensus, + split, + minimum_size, + AttentionCaptureProfile.BALANCED, + instance_recall=instance_recall, + geometry_recall=geometry_recall, + keep_only=keep_only, + combine_segs=combine, + matte_solidity=solidity, + edge_feather=feather, + evidence_mode=evidence_mode, + ) diff --git a/tests/test_attention_region_store.py b/tests/regional_generation/attention_regions/test_attention_region_store.py similarity index 100% rename from tests/test_attention_region_store.py rename to tests/regional_generation/attention_regions/test_attention_region_store.py diff --git a/tests/test_attention_region_support.py b/tests/regional_generation/attention_regions/test_attention_region_support.py similarity index 100% rename from tests/test_attention_region_support.py rename to tests/regional_generation/attention_regions/test_attention_region_support.py diff --git a/tests/test_attention_region_support_topology.py b/tests/regional_generation/attention_regions/test_attention_region_support_topology.py similarity index 100% rename from tests/test_attention_region_support_topology.py rename to tests/regional_generation/attention_regions/test_attention_region_support_topology.py diff --git a/tests/test_attention_region_tokens.py b/tests/regional_generation/attention_regions/test_attention_region_tokens.py similarity index 100% rename from tests/test_attention_region_tokens.py rename to tests/regional_generation/attention_regions/test_attention_region_tokens.py diff --git a/tests/test_attention_region_v3_nodes.py b/tests/regional_generation/attention_regions/test_attention_region_v3_nodes.py similarity index 100% rename from tests/test_attention_region_v3_nodes.py rename to tests/regional_generation/attention_regions/test_attention_region_v3_nodes.py diff --git a/tests/test_attention_spatial_projection.py b/tests/regional_generation/attention_regions/test_attention_spatial_projection.py similarity index 100% rename from tests/test_attention_spatial_projection.py rename to tests/regional_generation/attention_regions/test_attention_spatial_projection.py diff --git a/tests/regional_generation/regional/__init__.py b/tests/regional_generation/regional/__init__.py new file mode 100644 index 0000000..061c77e --- /dev/null +++ b/tests/regional_generation/regional/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation regional test behavior.""" diff --git a/tests/regional_generation/regional/support/__init__.py b/tests/regional_generation/regional/support/__init__.py new file mode 100644 index 0000000..1c4a063 --- /dev/null +++ b/tests/regional_generation/regional/support/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation regional support test behavior.""" diff --git a/tests/regional_attention_test_values.py b/tests/regional_generation/regional/support/regional_attention_test_values.py similarity index 100% rename from tests/regional_attention_test_values.py rename to tests/regional_generation/regional/support/regional_attention_test_values.py diff --git a/tests/regional_lora_test_values.py b/tests/regional_generation/regional/support/regional_lora_test_values.py similarity index 98% rename from tests/regional_lora_test_values.py rename to tests/regional_generation/regional/support/regional_lora_test_values.py index 7f286d8..6203f9c 100644 --- a/tests/regional_lora_test_values.py +++ b/tests/regional_generation/regional/support/regional_lora_test_values.py @@ -9,7 +9,6 @@ from __future__ import annotations from uuid import uuid4 import torch -from regional_attention_test_values import single_entry_regions from simple_syrup.domain.regional_attention import RegionalAttentionBranch from simple_syrup.domain.regional_attention_batch import ( @@ -53,6 +52,10 @@ from simple_syrup.runtime.regional_lora_schedule_resolution import ( RegionalLoraScheduleResolution, ) +from .regional_attention_test_values import ( + single_entry_regions, +) + def static_lora_schedule( composition: AnimaRegionalLoraComposition, diff --git a/tests/regional_model_capability_test_values.py b/tests/regional_generation/regional/support/regional_model_capability_test_values.py similarity index 100% rename from tests/regional_model_capability_test_values.py rename to tests/regional_generation/regional/support/regional_model_capability_test_values.py diff --git a/tests/test_global_hook_model_resolver.py b/tests/regional_generation/regional/test_global_hook_model_resolver.py similarity index 100% rename from tests/test_global_hook_model_resolver.py rename to tests/regional_generation/regional/test_global_hook_model_resolver.py diff --git a/tests/test_global_prompt_lora_proof.py b/tests/regional_generation/regional/test_global_prompt_lora_proof.py similarity index 100% rename from tests/test_global_prompt_lora_proof.py rename to tests/regional_generation/regional/test_global_prompt_lora_proof.py diff --git a/tests/test_global_regional_lora_image_validation.py b/tests/regional_generation/regional/test_global_regional_lora_image_validation.py similarity index 100% rename from tests/test_global_regional_lora_image_validation.py rename to tests/regional_generation/regional/test_global_regional_lora_image_validation.py diff --git a/tests/test_global_regional_lora_integration.py b/tests/regional_generation/regional/test_global_regional_lora_integration.py similarity index 100% rename from tests/test_global_regional_lora_integration.py rename to tests/regional_generation/regional/test_global_regional_lora_integration.py diff --git a/tests/test_global_style_character_proof.py b/tests/regional_generation/regional/test_global_style_character_proof.py similarity index 100% rename from tests/test_global_style_character_proof.py rename to tests/regional_generation/regional/test_global_style_character_proof.py diff --git a/tests/test_native_hooked_regional_conditioning.py b/tests/regional_generation/regional/test_native_hooked_regional_conditioning.py similarity index 100% rename from tests/test_native_hooked_regional_conditioning.py rename to tests/regional_generation/regional/test_native_hooked_regional_conditioning.py diff --git a/tests/test_regional_activation_batch_alignment.py b/tests/regional_generation/regional/test_regional_activation_batch_alignment.py similarity index 100% rename from tests/test_regional_activation_batch_alignment.py rename to tests/regional_generation/regional/test_regional_activation_batch_alignment.py diff --git a/tests/test_regional_activation_geometry.py b/tests/regional_generation/regional/test_regional_activation_geometry.py similarity index 100% rename from tests/test_regional_activation_geometry.py rename to tests/regional_generation/regional/test_regional_activation_geometry.py diff --git a/tests/test_regional_activation_geometry_providers.py b/tests/regional_generation/regional/test_regional_activation_geometry_providers.py similarity index 100% rename from tests/test_regional_activation_geometry_providers.py rename to tests/regional_generation/regional/test_regional_activation_geometry_providers.py diff --git a/tests/test_regional_activation_mask_projection.py b/tests/regional_generation/regional/test_regional_activation_mask_projection.py similarity index 100% rename from tests/test_regional_activation_mask_projection.py rename to tests/regional_generation/regional/test_regional_activation_mask_projection.py diff --git a/tests/test_regional_attention.py b/tests/regional_generation/regional/test_regional_attention.py similarity index 66% rename from tests/test_regional_attention.py rename to tests/regional_generation/regional/test_regional_attention.py index 078cfe5..408377a 100644 --- a/tests/test_regional_attention.py +++ b/tests/regional_generation/regional/test_regional_attention.py @@ -25,10 +25,6 @@ from simple_syrup.domain.raw_regional_attention import ( RawRegionalAttentionContext, build_raw_regional_attention_plan, ) -from simple_syrup.domain.regional_attention import RegionalAttentionBranch -from simple_syrup.domain.regional_attention_selection import ( - REGIONAL_ATTENTION_SELECTION_SERVICE, -) from simple_syrup.domain.regional_lora_plan import ( RegionalLoraAdapterIdentity, RegionalLoraAdapterPlan, @@ -246,125 +242,6 @@ def test_attention_plan_values_are_frozen() -> None: processed.region_index = 0 # type: ignore[misc] -@pytest.mark.parametrize( - ("selectors", "expected"), - [ - ([], []), - ([0], [RegionalAttentionBranch.POSITIVE]), - ([1], [RegionalAttentionBranch.NEGATIVE]), - ([0, 1], [RegionalAttentionBranch.POSITIVE, RegionalAttentionBranch.NEGATIVE]), - ([1, 0], [RegionalAttentionBranch.NEGATIVE, RegionalAttentionBranch.POSITIVE]), - ( - [0, 0, 1, 0, 1, 1], - [ - RegionalAttentionBranch.POSITIVE, - RegionalAttentionBranch.POSITIVE, - RegionalAttentionBranch.NEGATIVE, - RegionalAttentionBranch.POSITIVE, - RegionalAttentionBranch.NEGATIVE, - RegionalAttentionBranch.NEGATIVE, - ], - ), - ], -) -def test_processed_chunk_selection_follows_exact_comfy_order( - selectors: list[int], - expected: list[RegionalAttentionBranch], -) -> None: - """Support CFG-disabled, reversed, repeated, and interleaved chunk sequences.""" - - plan = _processed_plan() - - chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=selectors, - conditioning_uuids=_uuids_for_selectors(plan, selectors), - sigma=0.5, - ) - - assert [chunk.chunk_index for chunk in chunks] == list(range(len(selectors))) - assert [chunk.branch for chunk in chunks] == expected - assert all( - len(chunk.regional_entries) == plan.mask_bank.region_count for chunk in chunks - ) - - -@pytest.mark.parametrize( - ("selectors", "error", "message"), - [ - (object(), TypeError, "list or tuple"), - ([True], TypeError, "integer 0 or 1"), - (["0"], TypeError, "integer 0 or 1"), - ([2], ValueError, "observed 2"), - ([-1], ValueError, "observed -1"), - ], -) -def test_processed_chunk_selection_rejects_invalid_selectors( - selectors: object, - error: type[Exception], - message: str, -) -> None: - """Fail closed instead of guessing conditional branch order.""" - - with pytest.raises(error, match=message): - REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - _processed_plan(), - cond_or_uncond=selectors, - conditioning_uuids=( - [uuid4()] * len(selectors) if isinstance(selectors, list) else [] - ), - sigma=0.5, - ) - - -def test_processed_chunk_selection_preserves_all_active_base_entries() -> None: - """Map repeated branch chunks to simultaneous base entries in declared order.""" - - plan = _processed_plan() - positive_base = ProcessedRegionalAttentionContext( - 0, - None, - ( - _entry(0, torch.full((1, 2, 3), 1.0), 0.25), - _entry(1, torch.full((1, 2, 3), 5.0), 0.75), - ), - ) - plan = ProcessedRegionalAttentionPlan( - ProcessedRegionalAttentionBranch( - positive_base, - plan.positive.regional_contexts, - ), - plan.negative, - plan.mask_bank, - plan.lora_plan, - ) - - chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0, 0, 1], - conditioning_uuids=[ - positive_base.entries[0].uuid, - positive_base.entries[1].uuid, - plan.negative.base_context.entries[0].uuid, - ], - sigma=0.5, - ) - - assert [chunk.base_entry.entry_index for chunk in chunks] == [0, 1, 0] - assert [chunk.base_entry.strength for chunk in chunks] == [0.25, 0.75, 1.0] - with pytest.raises(ValueError, match="does not identify exactly one"): - REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0, 0, 0], - conditioning_uuids=[ - positive_base.entries[0].uuid, - positive_base.entries[1].uuid, - uuid4(), - ], - sigma=0.5, - ) - - def _processed( conditioning_index: int, region_index: int | None, @@ -396,39 +273,6 @@ def _entry( ) -def _processed_plan() -> ProcessedRegionalAttentionPlan: - """Build distinguishable positive and negative processed banks.""" - - return ProcessedRegionalAttentionPlan( - ProcessedRegionalAttentionBranch( - _processed(0, None, value=1.0), - (_processed(1, 0, value=2.0),), - ), - ProcessedRegionalAttentionBranch( - _processed(0, None, value=-1.0), - (_processed(1, 0, value=-2.0),), - ), - _mask_bank(1), - RegionalLoraPlan(()), - ) - - -def _uuids_for_selectors( - plan: ProcessedRegionalAttentionPlan, - selectors: list[int], -) -> list[object]: - """Return matching static UUIDs for arbitrary characterized chunk order.""" - - return [ - ( - plan.positive.base_context.entries[0].uuid - if selector == 0 - else plan.negative.base_context.entries[0].uuid - ) - for selector in selectors - ] - - def _mask_bank(region_count: int) -> RegionalMaskBank: """Build one canonical bank with distinct tensor storage.""" diff --git a/tests/test_regional_attention_active_sequence_alignment.py b/tests/regional_generation/regional/test_regional_attention_active_sequence_alignment.py similarity index 100% rename from tests/test_regional_attention_active_sequence_alignment.py rename to tests/regional_generation/regional/test_regional_attention_active_sequence_alignment.py diff --git a/tests/test_regional_attention_batching.py b/tests/regional_generation/regional/test_regional_attention_batching.py similarity index 100% rename from tests/test_regional_attention_batching.py rename to tests/regional_generation/regional/test_regional_attention_batching.py diff --git a/tests/test_regional_attention_diagnostics.py b/tests/regional_generation/regional/test_regional_attention_diagnostics.py similarity index 100% rename from tests/test_regional_attention_diagnostics.py rename to tests/regional_generation/regional/test_regional_attention_diagnostics.py diff --git a/tests/test_regional_attention_execution_context.py b/tests/regional_generation/regional/test_regional_attention_execution_context.py similarity index 100% rename from tests/test_regional_attention_execution_context.py rename to tests/regional_generation/regional/test_regional_attention_execution_context.py diff --git a/tests/test_regional_attention_model_call.py b/tests/regional_generation/regional/test_regional_attention_model_call.py similarity index 100% rename from tests/test_regional_attention_model_call.py rename to tests/regional_generation/regional/test_regional_attention_model_call.py diff --git a/tests/test_regional_attention_model_call_values.py b/tests/regional_generation/regional/test_regional_attention_model_call_values.py similarity index 100% rename from tests/test_regional_attention_model_call_values.py rename to tests/regional_generation/regional/test_regional_attention_model_call_values.py diff --git a/tests/test_regional_attention_query_masks.py b/tests/regional_generation/regional/test_regional_attention_query_masks.py similarity index 100% rename from tests/test_regional_attention_query_masks.py rename to tests/regional_generation/regional/test_regional_attention_query_masks.py diff --git a/tests/regional_generation/regional/test_regional_attention_selection.py b/tests/regional_generation/regional/test_regional_attention_selection.py new file mode 100644 index 0000000..5d0309f --- /dev/null +++ b/tests/regional_generation/regional/test_regional_attention_selection.py @@ -0,0 +1,215 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Prove processed regional-attention chunk selection contracts.""" + +from __future__ import annotations + +from uuid import uuid4 + +import pytest +import torch + +from simple_syrup.domain.conditioning_schedule import ( + UNBOUNDED_CONDITIONING_SCHEDULE, +) +from simple_syrup.domain.processed_regional_attention import ( + ProcessedRegionalAttentionBranch, + ProcessedRegionalAttentionContext, + ProcessedRegionalAttentionEntry, + ProcessedRegionalAttentionPlan, +) +from simple_syrup.domain.regional_attention import RegionalAttentionBranch +from simple_syrup.domain.regional_attention_selection import ( + REGIONAL_ATTENTION_SELECTION_SERVICE, +) +from simple_syrup.domain.regional_lora_plan import RegionalLoraPlan +from simple_syrup.domain.regional_mask_bank import RegionalMaskBank + + +@pytest.mark.parametrize( + ("selectors", "expected"), + [ + ([], []), + ([0], [RegionalAttentionBranch.POSITIVE]), + ([1], [RegionalAttentionBranch.NEGATIVE]), + ([0, 1], [RegionalAttentionBranch.POSITIVE, RegionalAttentionBranch.NEGATIVE]), + ([1, 0], [RegionalAttentionBranch.NEGATIVE, RegionalAttentionBranch.POSITIVE]), + ( + [0, 0, 1, 0, 1, 1], + [ + RegionalAttentionBranch.POSITIVE, + RegionalAttentionBranch.POSITIVE, + RegionalAttentionBranch.NEGATIVE, + RegionalAttentionBranch.POSITIVE, + RegionalAttentionBranch.NEGATIVE, + RegionalAttentionBranch.NEGATIVE, + ], + ), + ], +) +def test_processed_chunk_selection_follows_exact_comfy_order( + selectors: list[int], + expected: list[RegionalAttentionBranch], +) -> None: + """Support CFG-disabled, reversed, repeated, and interleaved chunk sequences.""" + + plan = _processed_plan() + chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( + plan, + cond_or_uncond=selectors, + conditioning_uuids=_uuids_for_selectors(plan, selectors), + sigma=0.5, + ) + assert [chunk.chunk_index for chunk in chunks] == list(range(len(selectors))) + assert [chunk.branch for chunk in chunks] == expected + assert all( + len(chunk.regional_entries) == plan.mask_bank.region_count for chunk in chunks + ) + + +@pytest.mark.parametrize( + ("selectors", "error", "message"), + [ + (object(), TypeError, "list or tuple"), + ([True], TypeError, "integer 0 or 1"), + (["0"], TypeError, "integer 0 or 1"), + ([2], ValueError, "observed 2"), + ([-1], ValueError, "observed -1"), + ], +) +def test_processed_chunk_selection_rejects_invalid_selectors( + selectors: object, + error: type[Exception], + message: str, +) -> None: + """Fail closed instead of guessing conditional branch order.""" + + with pytest.raises(error, match=message): + REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( + _processed_plan(), + cond_or_uncond=selectors, + conditioning_uuids=( + [uuid4()] * len(selectors) if isinstance(selectors, list) else [] + ), + sigma=0.5, + ) + + +def test_processed_chunk_selection_preserves_all_active_base_entries() -> None: + """Map repeated branch chunks to simultaneous base entries in declared order.""" + + plan = _processed_plan() + positive_base = ProcessedRegionalAttentionContext( + 0, + None, + ( + _entry(0, torch.full((1, 2, 3), 1.0), 0.25), + _entry(1, torch.full((1, 2, 3), 5.0), 0.75), + ), + ) + plan = ProcessedRegionalAttentionPlan( + ProcessedRegionalAttentionBranch( + positive_base, + plan.positive.regional_contexts, + ), + plan.negative, + plan.mask_bank, + plan.lora_plan, + ) + chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( + plan, + cond_or_uncond=[0, 0, 1], + conditioning_uuids=[ + positive_base.entries[0].uuid, + positive_base.entries[1].uuid, + plan.negative.base_context.entries[0].uuid, + ], + sigma=0.5, + ) + assert [chunk.base_entry.entry_index for chunk in chunks] == [0, 1, 0] + assert [chunk.base_entry.strength for chunk in chunks] == [0.25, 0.75, 1.0] + with pytest.raises(ValueError, match="does not identify exactly one"): + REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( + plan, + cond_or_uncond=[0, 0, 0], + conditioning_uuids=[ + positive_base.entries[0].uuid, + positive_base.entries[1].uuid, + uuid4(), + ], + sigma=0.5, + ) + + +def _processed( + conditioning_index: int, + region_index: int | None, + *, + value: float, +) -> ProcessedRegionalAttentionContext: + """Build one small valid processed context.""" + + return ProcessedRegionalAttentionContext( + conditioning_index, + region_index, + (_entry(0, torch.full((1, 2, 3), value), 1.0),), + ) + + +def _entry( + entry_index: int, + tensor: torch.Tensor, + strength: float, +) -> ProcessedRegionalAttentionEntry: + """Build one unscheduled processed entry with an explicit Comfy identity.""" + + return ProcessedRegionalAttentionEntry( + entry_index, + uuid4(), + UNBOUNDED_CONDITIONING_SCHEDULE, + tensor, + strength, + ) + + +def _processed_plan() -> ProcessedRegionalAttentionPlan: + """Build distinguishable positive and negative processed banks.""" + + return ProcessedRegionalAttentionPlan( + ProcessedRegionalAttentionBranch( + _processed(0, None, value=1.0), + (_processed(1, 0, value=2.0),), + ), + ProcessedRegionalAttentionBranch( + _processed(0, None, value=-1.0), + (_processed(1, 0, value=-2.0),), + ), + _mask_bank(1), + RegionalLoraPlan(()), + ) + + +def _uuids_for_selectors( + plan: ProcessedRegionalAttentionPlan, + selectors: list[int], +) -> list[object]: + """Return matching static UUIDs for arbitrary characterized chunk order.""" + + return [ + ( + plan.positive.base_context.entries[0].uuid + if selector == 0 + else plan.negative.base_context.entries[0].uuid + ) + for selector in selectors + ] + + +def _mask_bank(region_count: int) -> RegionalMaskBank: + """Build one canonical bank with distinct tensor storage.""" + + planning = torch.zeros((region_count, 4, 5), dtype=torch.float32) + conditioning = torch.zeros_like(planning) + return RegionalMaskBank(planning, conditioning, 5, 4) diff --git a/tests/test_regional_attention_sequence_alignment.py b/tests/regional_generation/regional/test_regional_attention_sequence_alignment.py similarity index 100% rename from tests/test_regional_attention_sequence_alignment.py rename to tests/regional_generation/regional/test_regional_attention_sequence_alignment.py diff --git a/tests/test_regional_attention_weights.py b/tests/regional_generation/regional/test_regional_attention_weights.py similarity index 100% rename from tests/test_regional_attention_weights.py rename to tests/regional_generation/regional/test_regional_attention_weights.py diff --git a/tests/test_regional_attention_weights_comfy_equivalence.py b/tests/regional_generation/regional/test_regional_attention_weights_comfy_equivalence.py similarity index 100% rename from tests/test_regional_attention_weights_comfy_equivalence.py rename to tests/regional_generation/regional/test_regional_attention_weights_comfy_equivalence.py diff --git a/tests/test_regional_cache_execution_interop.py b/tests/regional_generation/regional/test_regional_cache_execution_interop.py similarity index 100% rename from tests/test_regional_cache_execution_interop.py rename to tests/regional_generation/regional/test_regional_cache_execution_interop.py diff --git a/tests/test_regional_capability_admission_service.py b/tests/regional_generation/regional/test_regional_capability_admission_service.py similarity index 100% rename from tests/test_regional_capability_admission_service.py rename to tests/regional_generation/regional/test_regional_capability_admission_service.py diff --git a/tests/test_regional_conditioning_companion.py b/tests/regional_generation/regional/test_regional_conditioning_companion.py similarity index 100% rename from tests/test_regional_conditioning_companion.py rename to tests/regional_generation/regional/test_regional_conditioning_companion.py diff --git a/tests/test_regional_conditioning_output.py b/tests/regional_generation/regional/test_regional_conditioning_output.py similarity index 100% rename from tests/test_regional_conditioning_output.py rename to tests/regional_generation/regional/test_regional_conditioning_output.py diff --git a/tests/test_regional_conditioning_service.py b/tests/regional_generation/regional/test_regional_conditioning_service.py similarity index 100% rename from tests/test_regional_conditioning_service.py rename to tests/regional_generation/regional/test_regional_conditioning_service.py diff --git a/tests/test_regional_context_validation.py b/tests/regional_generation/regional/test_regional_context_validation.py similarity index 100% rename from tests/test_regional_context_validation.py rename to tests/regional_generation/regional/test_regional_context_validation.py diff --git a/tests/test_regional_convolution_execution.py b/tests/regional_generation/regional/test_regional_convolution_execution.py similarity index 100% rename from tests/test_regional_convolution_execution.py rename to tests/regional_generation/regional/test_regional_convolution_execution.py diff --git a/tests/test_regional_convolution_execution_plan.py b/tests/regional_generation/regional/test_regional_convolution_execution_plan.py similarity index 100% rename from tests/test_regional_convolution_execution_plan.py rename to tests/regional_generation/regional/test_regional_convolution_execution_plan.py diff --git a/tests/test_regional_convolution_rank_geometry.py b/tests/regional_generation/regional/test_regional_convolution_rank_geometry.py similarity index 100% rename from tests/test_regional_convolution_rank_geometry.py rename to tests/regional_generation/regional/test_regional_convolution_rank_geometry.py diff --git a/tests/test_regional_detailing_domain.py b/tests/regional_generation/regional/test_regional_detailing_domain.py similarity index 100% rename from tests/test_regional_detailing_domain.py rename to tests/regional_generation/regional/test_regional_detailing_domain.py diff --git a/tests/test_regional_detailing_masks.py b/tests/regional_generation/regional/test_regional_detailing_masks.py similarity index 100% rename from tests/test_regional_detailing_masks.py rename to tests/regional_generation/regional/test_regional_detailing_masks.py diff --git a/tests/test_regional_diagnostics_capture_probe.py b/tests/regional_generation/regional/test_regional_diagnostics_capture_probe.py similarity index 100% rename from tests/test_regional_diagnostics_capture_probe.py rename to tests/regional_generation/regional/test_regional_diagnostics_capture_probe.py diff --git a/tests/test_regional_features.py b/tests/regional_generation/regional/test_regional_features.py similarity index 100% rename from tests/test_regional_features.py rename to tests/regional_generation/regional/test_regional_features.py diff --git a/tests/test_regional_hook_fixture_probe.py b/tests/regional_generation/regional/test_regional_hook_fixture_probe.py similarity index 100% rename from tests/test_regional_hook_fixture_probe.py rename to tests/regional_generation/regional/test_regional_hook_fixture_probe.py diff --git a/tests/test_regional_host_linear_parameters.py b/tests/regional_generation/regional/test_regional_host_linear_parameters.py similarity index 100% rename from tests/test_regional_host_linear_parameters.py rename to tests/regional_generation/regional/test_regional_host_linear_parameters.py diff --git a/tests/test_regional_host_operation_backing.py b/tests/regional_generation/regional/test_regional_host_operation_backing.py similarity index 100% rename from tests/test_regional_host_operation_backing.py rename to tests/regional_generation/regional/test_regional_host_operation_backing.py diff --git a/tests/test_regional_ksampler_v3_nodes.py b/tests/regional_generation/regional/test_regional_ksampler_v3_nodes.py similarity index 100% rename from tests/test_regional_ksampler_v3_nodes.py rename to tests/regional_generation/regional/test_regional_ksampler_v3_nodes.py diff --git a/tests/test_regional_linear_execution.py b/tests/regional_generation/regional/test_regional_linear_execution.py similarity index 100% rename from tests/test_regional_linear_execution.py rename to tests/regional_generation/regional/test_regional_linear_execution.py diff --git a/tests/test_regional_linear_execution_plan.py b/tests/regional_generation/regional/test_regional_linear_execution_plan.py similarity index 100% rename from tests/test_regional_linear_execution_plan.py rename to tests/regional_generation/regional/test_regional_linear_execution_plan.py diff --git a/tests/test_regional_linear_mapped_projection_plan.py b/tests/regional_generation/regional/test_regional_linear_mapped_projection_plan.py similarity index 100% rename from tests/test_regional_linear_mapped_projection_plan.py rename to tests/regional_generation/regional/test_regional_linear_mapped_projection_plan.py diff --git a/tests/test_regional_lora_active_support.py b/tests/regional_generation/regional/test_regional_lora_active_support.py similarity index 100% rename from tests/test_regional_lora_active_support.py rename to tests/regional_generation/regional/test_regional_lora_active_support.py diff --git a/tests/test_regional_lora_compatible_rank_projection.py b/tests/regional_generation/regional/test_regional_lora_compatible_rank_projection.py similarity index 100% rename from tests/test_regional_lora_compatible_rank_projection.py rename to tests/regional_generation/regional/test_regional_lora_compatible_rank_projection.py diff --git a/tests/test_regional_lora_conditioning_adapter.py b/tests/regional_generation/regional/test_regional_lora_conditioning_adapter.py similarity index 100% rename from tests/test_regional_lora_conditioning_adapter.py rename to tests/regional_generation/regional/test_regional_lora_conditioning_adapter.py diff --git a/tests/test_regional_lora_delta_execution.py b/tests/regional_generation/regional/test_regional_lora_delta_execution.py similarity index 100% rename from tests/test_regional_lora_delta_execution.py rename to tests/regional_generation/regional/test_regional_lora_delta_execution.py diff --git a/tests/test_regional_lora_execution_cache.py b/tests/regional_generation/regional/test_regional_lora_execution_cache.py similarity index 100% rename from tests/test_regional_lora_execution_cache.py rename to tests/regional_generation/regional/test_regional_lora_execution_cache.py diff --git a/tests/test_regional_lora_fused_active_accumulation.py b/tests/regional_generation/regional/test_regional_lora_fused_active_accumulation.py similarity index 100% rename from tests/test_regional_lora_fused_active_accumulation.py rename to tests/regional_generation/regional/test_regional_lora_fused_active_accumulation.py diff --git a/tests/test_regional_lora_fused_multiplier_transport.py b/tests/regional_generation/regional/test_regional_lora_fused_multiplier_transport.py similarity index 100% rename from tests/test_regional_lora_fused_multiplier_transport.py rename to tests/regional_generation/regional/test_regional_lora_fused_multiplier_transport.py diff --git a/tests/test_regional_lora_hook_identity.py b/tests/regional_generation/regional/test_regional_lora_hook_identity.py similarity index 100% rename from tests/test_regional_lora_hook_identity.py rename to tests/regional_generation/regional/test_regional_lora_hook_identity.py diff --git a/tests/test_regional_lora_hooks.py b/tests/regional_generation/regional/test_regional_lora_hooks.py similarity index 100% rename from tests/test_regional_lora_hooks.py rename to tests/regional_generation/regional/test_regional_lora_hooks.py diff --git a/tests/test_regional_lora_host_payload.py b/tests/regional_generation/regional/test_regional_lora_host_payload.py similarity index 100% rename from tests/test_regional_lora_host_payload.py rename to tests/regional_generation/regional/test_regional_lora_host_payload.py diff --git a/tests/test_regional_lora_host_payloads.py b/tests/regional_generation/regional/test_regional_lora_host_payloads.py similarity index 100% rename from tests/test_regional_lora_host_payloads.py rename to tests/regional_generation/regional/test_regional_lora_host_payloads.py diff --git a/tests/test_regional_lora_plan.py b/tests/regional_generation/regional/test_regional_lora_plan.py similarity index 100% rename from tests/test_regional_lora_plan.py rename to tests/regional_generation/regional/test_regional_lora_plan.py diff --git a/tests/test_regional_lora_schedule_comfy_characterization.py b/tests/regional_generation/regional/test_regional_lora_schedule_comfy_characterization.py similarity index 100% rename from tests/test_regional_lora_schedule_comfy_characterization.py rename to tests/regional_generation/regional/test_regional_lora_schedule_comfy_characterization.py diff --git a/tests/test_regional_lora_schedule_resolution.py b/tests/regional_generation/regional/test_regional_lora_schedule_resolution.py similarity index 100% rename from tests/test_regional_lora_schedule_resolution.py rename to tests/regional_generation/regional/test_regional_lora_schedule_resolution.py diff --git a/tests/test_regional_lora_single_active_projection.py b/tests/regional_generation/regional/test_regional_lora_single_active_projection.py similarity index 100% rename from tests/test_regional_lora_single_active_projection.py rename to tests/regional_generation/regional/test_regional_lora_single_active_projection.py diff --git a/tests/test_regional_lora_target_binder.py b/tests/regional_generation/regional/test_regional_lora_target_binder.py similarity index 100% rename from tests/test_regional_lora_target_binder.py rename to tests/regional_generation/regional/test_regional_lora_target_binder.py diff --git a/tests/test_regional_lora_target_shape_admission.py b/tests/regional_generation/regional/test_regional_lora_target_shape_admission.py similarity index 100% rename from tests/test_regional_lora_target_shape_admission.py rename to tests/regional_generation/regional/test_regional_lora_target_shape_admission.py diff --git a/tests/test_regional_mask_activation.py b/tests/regional_generation/regional/test_regional_mask_activation.py similarity index 100% rename from tests/test_regional_mask_activation.py rename to tests/regional_generation/regional/test_regional_mask_activation.py diff --git a/tests/test_regional_mask_bank.py b/tests/regional_generation/regional/test_regional_mask_bank.py similarity index 100% rename from tests/test_regional_mask_bank.py rename to tests/regional_generation/regional/test_regional_mask_bank.py diff --git a/tests/test_regional_mask_projection.py b/tests/regional_generation/regional/test_regional_mask_projection.py similarity index 100% rename from tests/test_regional_mask_projection.py rename to tests/regional_generation/regional/test_regional_mask_projection.py diff --git a/tests/test_regional_model_capability_registry.py b/tests/regional_generation/regional/test_regional_model_capability_registry.py similarity index 98% rename from tests/test_regional_model_capability_registry.py rename to tests/regional_generation/regional/test_regional_model_capability_registry.py index e2895f4..f871326 100644 --- a/tests/test_regional_model_capability_registry.py +++ b/tests/regional_generation/regional/test_regional_model_capability_registry.py @@ -12,18 +12,19 @@ import comfy.model_patcher import pytest from comfy.ldm.anima.model import Anima as AnimaDiffusionModel from comfy.ldm.cosmos.predict2 import MiniTrainDIT -from regional_model_capability_test_values import ( - AlternateImageLatent, - empty_module, - patcher, - standard_unet_graph, -) from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily from simple_syrup.runtime.regional_model_capabilities import ( RegionalModelCapabilityRegistry, ) +from .support.regional_model_capability_test_values import ( + AlternateImageLatent, + empty_module, + patcher, + standard_unet_graph, +) + def test_registry_routes_one_capability_equivalent_standard_unet() -> None: """Select standard execution from live graph evidence rather than type pairs.""" diff --git a/tests/test_regional_model_capability_values.py b/tests/regional_generation/regional/test_regional_model_capability_values.py similarity index 100% rename from tests/test_regional_model_capability_values.py rename to tests/regional_generation/regional/test_regional_model_capability_values.py diff --git a/tests/test_regional_model_hook_selection.py b/tests/regional_generation/regional/test_regional_model_hook_selection.py similarity index 100% rename from tests/test_regional_model_hook_selection.py rename to tests/regional_generation/regional/test_regional_model_hook_selection.py diff --git a/tests/test_regional_model_modifier_characterization.py b/tests/regional_generation/regional/test_regional_model_modifier_characterization.py similarity index 98% rename from tests/test_regional_model_modifier_characterization.py rename to tests/regional_generation/regional/test_regional_model_modifier_characterization.py index 2750603..a88bee1 100644 --- a/tests/test_regional_model_modifier_characterization.py +++ b/tests/regional_generation/regional/test_regional_model_modifier_characterization.py @@ -8,7 +8,6 @@ from __future__ import annotations import importlib.util from collections.abc import Callable -from pathlib import Path from types import ModuleType, SimpleNamespace from typing import cast from uuid import uuid4 @@ -28,13 +27,14 @@ from comfy_extras.nodes_easycache import ( # type: ignore[import-not-found] easycache_sample_wrapper, lazycache_predict_noise_wrapper, ) +from support.repository import CUSTOM_NODES_ROOT from simple_syrup.runtime.model_patcher_mutations import ( ModelDiffusionWrapperMutation, ) from simple_syrup.runtime.patcher_lifecycle import PATCHER_LIFECYCLE -CUSTOM_NODES_ROOT = Path(__file__).resolve().parents[2] +CUSTOM_NODES_ROOT = CUSTOM_NODES_ROOT class _CacheFixtureModel(torch.nn.Module): diff --git a/tests/test_regional_model_patch_interop.py b/tests/regional_generation/regional/test_regional_model_patch_interop.py similarity index 100% rename from tests/test_regional_model_patch_interop.py rename to tests/regional_generation/regional/test_regional_model_patch_interop.py diff --git a/tests/test_regional_model_patch_stack.py b/tests/regional_generation/regional/test_regional_model_patch_stack.py similarity index 100% rename from tests/test_regional_model_patch_stack.py rename to tests/regional_generation/regional/test_regional_model_patch_stack.py diff --git a/tests/regional_generation/regional/test_regional_multidiffusion_conditioning.py b/tests/regional_generation/regional/test_regional_multidiffusion_conditioning.py new file mode 100644 index 0000000..9a4aaa4 --- /dev/null +++ b/tests/regional_generation/regional/test_regional_multidiffusion_conditioning.py @@ -0,0 +1,379 @@ +# 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 regional MultiDiffusion sampling runtime.""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any, cast + +import pytest +import torch + +from simple_syrup.domain.regional_detailing import LatentBox, LatentRegion +from simple_syrup.runtime import ( + regional_multidiffusion_prediction, + regional_multidiffusion_sampling, +) + +comfy_sample = regional_multidiffusion_sampling._comfy_sample() +comfy_utils = regional_multidiffusion_sampling._comfy_utils() +latent_preview = regional_multidiffusion_sampling._latent_preview() + + +class FakeModel: + """Provide the ModelPatcher methods used by the runtime.""" + + def __init__( + self, + model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, + ) -> None: + """Create a fake model patcher.""" + + self.load_device = torch.device("cpu") + 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.""" + + 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.""" + + 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.""" + + assert name == "model_sampling" + return self.model_sampling + + +class FakeSampler: + """Represent a resolved sampler in tests.""" + + def sample(self, *args: object, **kwargs: object) -> object: + """Provide ComfyUI's sampler protocol.""" + + del args, kwargs + return None + + +def test_raw_region_conditioning_is_converted_before_calc_cond_batch( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Raw CONDITIONING entries are converted before Comfy calc-cond-batch calls.""" + + calls: list[list[list[dict[str, Any]] | None]] = [] + + def calc_cond_batch( + _model: object, + conds: list[list[dict[str, Any]] | None], + x_in: torch.Tensor, + _timestep: torch.Tensor, + _model_options: dict[str, Any], + ) -> list[torch.Tensor]: + """Return constants while asserting sampler-ready conditioning shape.""" + + calls.append(conds) + first_cond = conds[0] + assert first_cond is not None + assert isinstance(first_cond[0], dict) + assert "model_conds" in first_cond[0] + value = 5.0 if "cross_attn" in first_cond[0] else 1.0 + return [torch.ones_like(x_in) * value, torch.zeros_like(x_in)] + + fake_samplers = SimpleNamespace( + calc_cond_batch=calc_cond_batch, + resolve_areas_and_cond_masks_multidim=lambda *_args: None, + calculate_start_end_timesteps=lambda *_args: None, + ) + monkeypatch.setattr( + regional_multidiffusion_prediction, + "_comfy_samplers", + lambda: fake_samplers, + ) + raw_region_positive = [[torch.ones((1, 1, 1)), {}]] + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel(), + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=(_region(0, 0, 4, 4, raw_region_positive),), + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": [ + [{"model_conds": {}, "uuid": object()}], + [{"model_conds": {}, "uuid": object()}], + ], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert len(calls) == 2 + assert "cross_attn" in cast(list[dict[str, Any]], calls[1][0])[0] + assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 5.0) + + +def test_global_prompt_weight_blends_full_region_with_global_prediction() -> None: + """Covered pixels keep the configured global positive prediction share.""" + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Return fallback or regional constants.""" + + x = cast(torch.Tensor, args["input"]) + cond = _condition_name(cast(list[object], args["conds"])[0]) + value = 10.0 if cond == "global" else 20.0 + return [torch.ones_like(x) * value, torch.zeros_like(x)] + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=(_region(0, 0, 4, 4, "regional"),), + global_prompt_weight=0.25, + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 17.5) + + +def test_global_prompt_weight_keeps_partial_mask_coverage() -> None: + """Soft masks scale regional influence before global/regional weighting.""" + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Return fallback or regional constants.""" + + x = cast(torch.Tensor, args["input"]) + cond = _condition_name(cast(list[object], args["conds"])[0]) + value = 10.0 if cond == "global" else 20.0 + return [torch.ones_like(x) * value, torch.zeros_like(x)] + + region = LatentRegion( + index=0, + label="soft", + latent_box=LatentBox(0, 0, 4, 4), + latent_mask=torch.ones((4, 4)) * 0.5, + positive=_raw_conditioning("regional"), + ) + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=(region,), + global_prompt_weight=0.25, + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 13.75) + + +def test_overlapping_regions_normalize_before_global_prompt_weight_blend() -> None: + """Overlaps average region predictions before applying global weight.""" + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Return different constants for each condition.""" + + x = cast(torch.Tensor, args["input"]) + cond = _condition_name(cast(list[object], args["conds"])[0]) + values = {"global": 10.0, "first": 20.0, "second": 40.0} + return [torch.ones_like(x) * values[cond], torch.zeros_like(x)] + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=( + _region(0, 0, 4, 4, "first"), + _region(0, 0, 4, 4, "second"), + ), + global_prompt_weight=0.25, + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 25.0) + + +def test_negative_conditioning_is_reused_for_region_unconditional_path() -> None: + """Regional calls keep the original negative conditioning.""" + + regional_conds: list[list[object]] = [] + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Record all conditioning lists.""" + + x = cast(torch.Tensor, args["input"]) + conds = cast(list[object], args["conds"]) + regional_conds.append(conds) + return [torch.ones_like(x), torch.ones_like(x) * 2.0] + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=(_region(0, 0, 4, 4, "regional"),), + ) + ) + + wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert [_condition_name(item) for item in regional_conds[1]] == [ + "regional", + "negative", + ] + + +def test_cfg_one_none_uncond_still_returns_two_entries() -> None: + """Comfy's CFG=1 optimization keeps the output list shape stable.""" + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Return one tensor per cond slot, even when uncond is None.""" + + x = cast(torch.Tensor, args["input"]) + conds = cast(list[object], args["conds"]) + value = 3.0 if _condition_name(conds[0]) == "regional" else 1.0 + return [torch.ones_like(x) * value, torch.zeros_like(x)] + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=(_region(0, 0, 4, 4, "regional"),), + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", None], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert len(output) == 2 + assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 3.0) + assert torch.allclose(output[1], torch.zeros((1, 1, 4, 4))) + + +def _region( + x: int, + y: int, + width: int, + height: int, + positive: object, + *, + latent_width: int = 4, + latent_height: int = 4, +) -> LatentRegion: + """Return one full-weight latent region.""" + + if isinstance(positive, str): + positive = _raw_conditioning(positive) + mask = torch.zeros((latent_height, latent_width), dtype=torch.float32) + mask[y : y + height, x : x + width] = 1.0 + return LatentRegion( + index=0, + label="region", + latent_box=LatentBox(x, y, width, height), + latent_mask=mask, + positive=positive, + ) + + +def _raw_conditioning(name: str) -> list[list[object]]: + """Return a raw Comfy CONDITIONING-like value with a visible test name.""" + + return [[torch.zeros((1, 1, 1), dtype=torch.float32), {"name": name}]] + + +def _condition_name(conditioning: object) -> str: + """Return the test-visible name from raw, processed, or sentinel conditioning.""" + + if isinstance(conditioning, str): + return conditioning + if isinstance(conditioning, list) and conditioning: + first = conditioning[0] + if isinstance(first, dict): + return str(first.get("name", "")) + if ( + isinstance(first, list | tuple) + and len(first) > 1 + and isinstance( + first[1], + dict, + ) + ): + return str(first[1].get("name", "")) + return "" diff --git a/tests/regional_generation/regional/test_regional_multidiffusion_sampling.py b/tests/regional_generation/regional/test_regional_multidiffusion_sampling.py new file mode 100644 index 0000000..078b8f3 --- /dev/null +++ b/tests/regional_generation/regional/test_regional_multidiffusion_sampling.py @@ -0,0 +1,505 @@ +# 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 regional MultiDiffusion sampling runtime.""" + +from __future__ import annotations + +from typing import Any, cast + +import pytest +import torch + +from simple_syrup.domain.regional_detailing import LatentBox, LatentRegion +from simple_syrup.domain.segs import CropRegion +from simple_syrup.runtime import ( + regional_multidiffusion_sampling, +) +from simple_syrup.runtime.detail_previews import DetailPreviewContext +from simple_syrup.runtime.regional_multidiffusion_prediction import ( + RegionalMultiDiffusionCalcCondBatch, +) + +comfy_sample = regional_multidiffusion_sampling._comfy_sample() +comfy_utils = regional_multidiffusion_sampling._comfy_utils() +latent_preview = regional_multidiffusion_sampling._latent_preview() + + +def _preview_context() -> DetailPreviewContext: + """Return a minimal regional detail preview context.""" + + return DetailPreviewContext( + image=torch.ones((1, 8, 8, 3), dtype=torch.float32), + work_region=CropRegion(2, 2, 6, 6), + work_mask=torch.ones((8, 8), dtype=torch.float32), + sampled_region=CropRegion(0, 0, 8, 8), + ) + + +def test_sampling_callback_uses_generic_preview_without_detail_context( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Standalone regional runtime calls keep generic latent previews.""" + + monkeypatch.setattr( + latent_preview, + "prepare_callback", + lambda _model, _steps: "generic callback", + ) + + assert ( + regional_multidiffusion_sampling._sampling_callback(FakeModel(), 4, None) + == "generic callback" + ) + + +def test_sampling_callback_uses_detail_preview_with_context( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Detailer regional sampling uses the shared detail preview callback.""" + + context = _preview_context() + calls: dict[str, object] = {} + + def fake_prepare_detail_preview_callback( + model: FakeModel, + steps: int, + preview_context: DetailPreviewContext, + ) -> str: + """Record detail preview callback preparation.""" + + calls["model"] = model + calls["steps"] = steps + calls["preview_context"] = preview_context + return "detail callback" + + monkeypatch.setattr( + regional_multidiffusion_sampling, + "prepare_detail_preview_callback", + fake_prepare_detail_preview_callback, + ) + + model = FakeModel() + assert ( + regional_multidiffusion_sampling._sampling_callback(model, 4, context) + == "detail callback" + ) + assert calls == {"model": model, "steps": 4, "preview_context": context} + + +class FakeModel: + """Provide the ModelPatcher methods used by the runtime.""" + + def __init__( + self, + model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, + ) -> None: + """Create a fake model patcher.""" + + self.load_device = torch.device("cpu") + 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.""" + + 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.""" + + 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.""" + + assert name == "model_sampling" + return self.model_sampling + + +class FakeSampler: + """Represent a resolved sampler in tests.""" + + def sample(self, *args: object, **kwargs: object) -> object: + """Provide ComfyUI's sampler protocol.""" + + del args, kwargs + return None + + +def test_clone_model_installs_regional_calc_cond_batch_wrapper() -> None: + """The runtime clones the model and installs a regional wrapper.""" + + 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),), + ) + ) + + assert wrapped_model is not model + assert isinstance( + wrapped_model.calc_wrapper, + RegionalMultiDiffusionCalcCondBatch, + ) + 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, + RegionalMultiDiffusionCalcCondBatch, + ) + + +def test_clone_model_rejects_non_callable_existing_calc_wrapper() -> None: + """Existing calc-cond-batch metadata must be callable.""" + + model = FakeModel({"sampler_calc_cond_batch_function": object()}) + + with pytest.raises(ValueError, match="Existing sampler_calc_cond_batch_function"): + 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),), + ) + + +def test_existing_calc_wrapper_is_composed_without_recursion() -> None: + """Fallback and regional calls delegate to the previous calc wrapper.""" + + calls: list[dict[str, Any]] = [] + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Return condition-specific constants and record options.""" + + calls.append(args) + assert args["model_options"].get("sampler_calc_cond_batch_function") is existing + x = cast(torch.Tensor, args["input"]) + conds = cast(list[object], args["conds"]) + value = 10.0 if _condition_name(conds[0]) == "global" else 20.0 + return [torch.ones_like(x) * value, torch.ones_like(x) * 2.0] + + model = FakeModel( + { + "sampler_calc_cond_batch_function": existing, + "model_function_wrapper": object(), + } + ) + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + model, + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=(_region(0, 0, 4, 4, "regional"),), + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert len(calls) == 2 + assert wrapped_model.model_options["model_function_wrapper"] is not None + assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 20.0) + assert torch.allclose(output[1], torch.ones((1, 1, 4, 4)) * 2.0) + + +def test_shape_mismatch_delegates_to_original_calc_path() -> None: + """Unexpected model input spatial shapes are delegated unchanged.""" + + calls = 0 + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Record fallback calls.""" + + nonlocal calls + calls += 1 + x = cast(torch.Tensor, args["input"]) + return [x + 5.0, x + 1.0] + + model = FakeModel({"sampler_calc_cond_batch_function": existing}) + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + model, + latent_width=8, + latent_height=8, + latent_ndim=4, + regions=(_region(0, 0, 4, 4, "regional", latent_width=8, latent_height=8),), + ) + ) + x = torch.zeros((1, 1, 4, 4)) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": x, + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert calls == 1 + assert torch.allclose(output[0], x + 5.0) + + +def test_region_crops_bchw_final_spatial_axes() -> None: + """Regional calls crop only final height and width axes.""" + + calls: list[torch.Tensor] = [] + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Record regional input crops.""" + + x = cast(torch.Tensor, args["input"]) + calls.append(x) + conds = cast(list[object], args["conds"]) + value = 1.0 if _condition_name(conds[0]) == "global" else 3.0 + return [torch.ones_like(x) * value, torch.zeros_like(x)] + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=8, + latent_height=4, + latent_ndim=4, + regions=(_region(0, 0, 8, 4, "regional", latent_width=8),), + ) + ) + x = torch.arange(32, dtype=torch.float32).reshape((1, 1, 4, 8)) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": x, + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert calls[1].shape == (1, 1, 4, 8) + assert torch.equal(calls[1], x) + assert torch.allclose(output[0], torch.ones((1, 1, 4, 8)) * 3.0) + + +def test_region_crops_singleton_depth_5d_final_spatial_axes() -> None: + """Anima-style regions crop only final height and width axes.""" + + calls: list[torch.Tensor] = [] + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Record regional input crops.""" + + x = cast(torch.Tensor, args["input"]) + calls.append(x) + conds = cast(list[object], args["conds"]) + value = 1.0 if _condition_name(conds[0]) == "global" else 4.0 + return [torch.ones_like(x) * value, torch.zeros_like(x)] + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=8, + latent_height=4, + latent_ndim=5, + regions=(_region(0, 0, 8, 4, "regional", latent_width=8),), + ) + ) + x = ( + torch.arange(128, dtype=torch.float32) + .reshape((1, 16, 1, 4, 2)) + .repeat(1, 1, 1, 1, 4) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": x, + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert calls[1].shape == (1, 16, 1, 4, 8) + assert torch.equal(calls[1], x) + assert output[0].shape == x.shape + + +def test_overlapping_regions_normalize_by_accumulated_weight() -> None: + """Overlapping regions are averaged before blending over fallback.""" + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Return different constants for each region positive.""" + + x = cast(torch.Tensor, args["input"]) + cond = _condition_name(cast(list[object], args["conds"])[0]) + values = {"global": 0.0, "first": 2.0, "second": 6.0} + return [torch.ones_like(x) * values[cond], torch.zeros_like(x)] + + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=6, + latent_height=4, + latent_ndim=4, + regions=( + _region(0, 0, 4, 4, "first", latent_width=6), + _region(2, 0, 4, 4, "second", latent_width=6), + ), + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": torch.zeros((1, 1, 4, 6)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert torch.allclose(output[0][:, :, :, :2], torch.ones((1, 1, 4, 2)) * 2.0) + assert torch.allclose(output[0][:, :, :, 2:4], torch.ones((1, 1, 4, 2)) * 4.0) + assert torch.allclose(output[0][:, :, :, 4:], torch.ones((1, 1, 4, 2)) * 6.0) + + +def test_partial_mask_blends_region_over_fallback() -> None: + """Feathered masks blend region predictions with fallback predictions.""" + + def existing(args: dict[str, Any]) -> list[torch.Tensor]: + """Return fallback or regional constants.""" + + x = cast(torch.Tensor, args["input"]) + cond = _condition_name(cast(list[object], args["conds"])[0]) + value = 10.0 if cond == "global" else 20.0 + return [torch.ones_like(x) * value, torch.zeros_like(x)] + + mask = torch.ones((4, 4)) * 0.25 + region = LatentRegion( + index=0, + label="soft", + latent_box=LatentBox(0, 0, 4, 4), + latent_mask=mask, + positive=_raw_conditioning("regional"), + ) + wrapped_model, _summary = ( + regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( + FakeModel({"sampler_calc_cond_batch_function": existing}), + latent_width=4, + latent_height=4, + latent_ndim=4, + regions=(region,), + ) + ) + + output = wrapped_model.calc_wrapper( + { + "conds": ["global", "negative"], + "input": torch.zeros((1, 1, 4, 4)), + "sigma": torch.tensor([1.0]), + "model": wrapped_model, + "model_options": wrapped_model.model_options, + } + ) + + assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 12.5) + + +def _region( + x: int, + y: int, + width: int, + height: int, + positive: object, + *, + latent_width: int = 4, + latent_height: int = 4, +) -> LatentRegion: + """Return one full-weight latent region.""" + + if isinstance(positive, str): + positive = _raw_conditioning(positive) + mask = torch.zeros((latent_height, latent_width), dtype=torch.float32) + mask[y : y + height, x : x + width] = 1.0 + return LatentRegion( + index=0, + label="region", + latent_box=LatentBox(x, y, width, height), + latent_mask=mask, + positive=positive, + ) + + +def _raw_conditioning(name: str) -> list[list[object]]: + """Return a raw Comfy CONDITIONING-like value with a visible test name.""" + + return [[torch.zeros((1, 1, 1), dtype=torch.float32), {"name": name}]] + + +def _condition_name(conditioning: object) -> str: + """Return the test-visible name from raw, processed, or sentinel conditioning.""" + + if isinstance(conditioning, str): + return conditioning + if isinstance(conditioning, list) and conditioning: + first = conditioning[0] + if isinstance(first, dict): + return str(first.get("name", "")) + if ( + isinstance(first, list | tuple) + and len(first) > 1 + and isinstance( + first[1], + dict, + ) + ): + return str(first[1].get("name", "")) + return "" diff --git a/tests/regional_generation/regional/test_regional_multidiffusion_sampling_runtime.py b/tests/regional_generation/regional/test_regional_multidiffusion_sampling_runtime.py new file mode 100644 index 0000000..3d171f3 --- /dev/null +++ b/tests/regional_generation/regional/test_regional_multidiffusion_sampling_runtime.py @@ -0,0 +1,343 @@ +# 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 regional MultiDiffusion sampling runtime.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.domain.regional_detailing import LatentBox, LatentRegion +from simple_syrup.runtime import ( + regional_multidiffusion_sampling, + sampling_samplers, + sampling_schedulers, +) + +comfy_sample = regional_multidiffusion_sampling._comfy_sample() +comfy_utils = regional_multidiffusion_sampling._comfy_utils() +latent_preview = regional_multidiffusion_sampling._latent_preview() + + +class FakeModel: + """Provide the ModelPatcher methods used by the runtime.""" + + def __init__( + self, + model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, + ) -> None: + """Create a fake model patcher.""" + + self.load_device = torch.device("cpu") + 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.""" + + 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.""" + + 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.""" + + assert name == "model_sampling" + return self.model_sampling + + +class FakeSampler: + """Represent a resolved sampler in tests.""" + + def sample(self, *args: object, **kwargs: object) -> object: + """Provide ComfyUI's sampler protocol.""" + + del args, kwargs + return None + + +def test_sample_rejects_unipc_before_sampler_resolution( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """UniPC sampler names fail before ComfyUI sampler lookup.""" + + def fail_resolve_sampler(_sampler_name: str) -> FakeSampler: + """Fail if sampler resolution is reached.""" + + raise AssertionError("UniPC rejection should happen before sampler resolution") + + monkeypatch.setattr(sampling_samplers, "resolve_sampler", fail_resolve_sampler) + + with pytest.raises(ValueError, match="not compatible with UniPC"): + regional_multidiffusion_sampling.sample_regional_multidiffusion( + model=FakeModel(), + seed=1, + steps=1, + cfg=1.0, + sampler_name="uni_pc", + scheduler="normal", + positive=[], + negative=[], + latent_image={"samples": torch.zeros((1, 4, 4, 4))}, + regions=(_region(0, 0, 4, 4, "regional"),), + denoise=1.0, + global_prompt_weight=0.0, + ) + + +def test_sample_rejects_unsupported_conditioning() -> None: + """Regional and ControlNet conditioning fail closed.""" + + with pytest.raises(ValueError, match="regional conditioning or ControlNet"): + regional_multidiffusion_sampling.sample_regional_multidiffusion( + model=FakeModel(), + seed=1, + steps=1, + cfg=1.0, + sampler_name="euler", + scheduler="normal", + positive=[{"area": (4, 4, 0, 0)}], + negative=[], + latent_image={"samples": torch.zeros((1, 4, 4, 4))}, + regions=(_region(0, 0, 4, 4, "regional"),), + denoise=1.0, + global_prompt_weight=0.0, + ) + + +def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Sampling mirrors KSampler flow while using a wrapped model clone.""" + + calls: dict[str, Any] = {} + model = FakeModel() + sampler = FakeSampler() + latent_samples = torch.zeros((1, 4, 4, 8), dtype=torch.float32) + fixed_noise = torch.ones_like(latent_samples) + fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) + sampled = torch.full_like(latent_samples, 0.25) + latent_image: dict[str, Any] = { + "samples": latent_samples, + "downscale_ratio_spacial": 2, + "kept": "value", + } + + 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=None: fixed_noise, + ) + monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) + + 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: object, + disable_pbar: bool, + seed: int, + ) -> torch.Tensor: + """Record sample_custom arguments.""" + + calls["sample_custom"] = { + "model": received_model, + "noise": noise, + "cfg": cfg, + "sampler": received_sampler, + "sigmas": sigmas, + "positive": positive, + "negative": negative, + "latent_image": latent_image, + "noise_mask": noise_mask, + "callback": callback, + "disable_pbar": disable_pbar, + "seed": seed, + } + return sampled + + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + + output = regional_multidiffusion_sampling.sample_regional_multidiffusion( + model=model, + seed=123, + steps=2, + cfg=7.0, + sampler_name="euler", + scheduler="normal", + positive=[{"model_conds": {}}], + negative=[{"model_conds": {}}], + latent_image=latent_image, + regions=(_region(0, 0, 4, 4, [{"model_conds": {}}], latent_width=8),), + denoise=1.0, + global_prompt_weight=0.25, + ) + + assert output is not latent_image + assert output["samples"] is sampled + assert output["kept"] == "value" + assert "downscale_ratio_spacial" not in output + assert calls["sample_custom"]["model"] is not model + assert calls["sample_custom"]["model"].calc_wrapper is not None + assert calls["sample_custom"]["sampler"] is sampler + assert calls["sample_custom"]["noise"] is fixed_noise + assert calls["sample_custom"]["disable_pbar"] is True + + +def test_sample_accepts_singleton_depth_5d_latent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Anima-style singleton-depth latents pass runtime validation.""" + + def fake_sample_custom( + _model: object, + _noise: torch.Tensor, + _cfg: float, + _sampler: object, + _sigmas: torch.Tensor, + _positive: object, + _negative: object, + latent_image: torch.Tensor, + **_kwargs: object, + ) -> torch.Tensor: + """Return a deterministic sampled latent for shape validation.""" + + return latent_image + 1.0 + + model = FakeModel() + sampler = FakeSampler() + latent_samples = torch.zeros((1, 16, 1, 4, 8), dtype=torch.float32) + fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) + + 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=None: torch.ones_like(samples), + ) + monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", True) + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + + output = regional_multidiffusion_sampling.sample_regional_multidiffusion( + model=model, + seed=123, + steps=2, + cfg=7.0, + sampler_name="euler", + scheduler="normal", + positive=[{"model_conds": {}}], + negative=[{"model_conds": {}}], + latent_image={"samples": latent_samples}, + regions=(_region(0, 0, 4, 4, [{"model_conds": {}}], latent_width=8),), + denoise=1.0, + global_prompt_weight=0.25, + ) + + assert torch.equal(output["samples"], latent_samples + 1.0) + + +def _region( + x: int, + y: int, + width: int, + height: int, + positive: object, + *, + latent_width: int = 4, + latent_height: int = 4, +) -> LatentRegion: + """Return one full-weight latent region.""" + + if isinstance(positive, str): + positive = _raw_conditioning(positive) + mask = torch.zeros((latent_height, latent_width), dtype=torch.float32) + mask[y : y + height, x : x + width] = 1.0 + return LatentRegion( + index=0, + label="region", + latent_box=LatentBox(x, y, width, height), + latent_mask=mask, + positive=positive, + ) + + +def _raw_conditioning(name: str) -> list[list[object]]: + """Return a raw Comfy CONDITIONING-like value with a visible test name.""" + + return [[torch.zeros((1, 1, 1), dtype=torch.float32), {"name": name}]] + + +def _condition_name(conditioning: object) -> str: + """Return the test-visible name from raw, processed, or sentinel conditioning.""" + + if isinstance(conditioning, str): + return conditioning + if isinstance(conditioning, list) and conditioning: + first = conditioning[0] + if isinstance(first, dict): + return str(first.get("name", "")) + if ( + isinstance(first, list | tuple) + and len(first) > 1 + and isinstance( + first[1], + dict, + ) + ): + return str(first[1].get("name", "")) + return "" diff --git a/tests/test_regional_operation_cache_lifecycle.py b/tests/regional_generation/regional/test_regional_operation_cache_lifecycle.py similarity index 100% rename from tests/test_regional_operation_cache_lifecycle.py rename to tests/regional_generation/regional/test_regional_operation_cache_lifecycle.py diff --git a/tests/test_regional_operation_call_scope.py b/tests/regional_generation/regional/test_regional_operation_call_scope.py similarity index 100% rename from tests/test_regional_operation_call_scope.py rename to tests/regional_generation/regional/test_regional_operation_call_scope.py diff --git a/tests/test_regional_operation_mask_resolution.py b/tests/regional_generation/regional/test_regional_operation_mask_resolution.py similarity index 100% rename from tests/test_regional_operation_mask_resolution.py rename to tests/regional_generation/regional/test_regional_operation_mask_resolution.py diff --git a/tests/test_regional_partitioned_linear_execution.py b/tests/regional_generation/regional/test_regional_partitioned_linear_execution.py similarity index 100% rename from tests/test_regional_partitioned_linear_execution.py rename to tests/regional_generation/regional/test_regional_partitioned_linear_execution.py diff --git a/tests/test_regional_patch_interop_evidence.py b/tests/regional_generation/regional/test_regional_patch_interop_evidence.py similarity index 100% rename from tests/test_regional_patch_interop_evidence.py rename to tests/regional_generation/regional/test_regional_patch_interop_evidence.py diff --git a/tests/test_regional_patch_interop_image_evidence.py b/tests/regional_generation/regional/test_regional_patch_interop_image_evidence.py similarity index 100% rename from tests/test_regional_patch_interop_image_evidence.py rename to tests/regional_generation/regional/test_regional_patch_interop_image_evidence.py diff --git a/tests/test_regional_patch_interop_results.py b/tests/regional_generation/regional/test_regional_patch_interop_results.py similarity index 100% rename from tests/test_regional_patch_interop_results.py rename to tests/regional_generation/regional/test_regional_patch_interop_results.py diff --git a/tests/test_regional_patch_interop_workflow.py b/tests/regional_generation/regional/test_regional_patch_interop_workflow.py similarity index 100% rename from tests/test_regional_patch_interop_workflow.py rename to tests/regional_generation/regional/test_regional_patch_interop_workflow.py diff --git a/tests/test_regional_prompt_masks.py b/tests/regional_generation/regional/test_regional_prompt_masks.py similarity index 100% rename from tests/test_regional_prompt_masks.py rename to tests/regional_generation/regional/test_regional_prompt_masks.py diff --git a/tests/test_regional_prompt_workflow.py b/tests/regional_generation/regional/test_regional_prompt_workflow.py similarity index 100% rename from tests/test_regional_prompt_workflow.py rename to tests/regional_generation/regional/test_regional_prompt_workflow.py diff --git a/tests/test_regional_prompting_domain.py b/tests/regional_generation/regional/test_regional_prompting_domain.py similarity index 100% rename from tests/test_regional_prompting_domain.py rename to tests/regional_generation/regional/test_regional_prompting_domain.py diff --git a/tests/test_regional_sampling_layout_characterization.py b/tests/regional_generation/regional/test_regional_sampling_layout_characterization.py similarity index 100% rename from tests/test_regional_sampling_layout_characterization.py rename to tests/regional_generation/regional/test_regional_sampling_layout_characterization.py diff --git a/tests/test_regional_sampling_preparation_service.py b/tests/regional_generation/regional/test_regional_sampling_preparation_service.py similarity index 100% rename from tests/test_regional_sampling_preparation_service.py rename to tests/regional_generation/regional/test_regional_sampling_preparation_service.py diff --git a/tests/test_regional_schedule_ownership.py b/tests/regional_generation/regional/test_regional_schedule_ownership.py similarity index 97% rename from tests/test_regional_schedule_ownership.py rename to tests/regional_generation/regional/test_regional_schedule_ownership.py index 1227c98..3892ecd 100644 --- a/tests/test_regional_schedule_ownership.py +++ b/tests/regional_generation/regional/test_regional_schedule_ownership.py @@ -7,16 +7,16 @@ from __future__ import annotations import ast -from pathlib import Path import pytest import torch +from support.repository import REPOSITORY_ROOT from simple_syrup.runtime.regional_attention_model_call_values import ( uniform_model_call_sigma, ) -_REPOSITORY_ROOT = Path(__file__).parents[1] +_REPOSITORY_ROOT = REPOSITORY_ROOT _RUNTIME_SELECTORS = ( "simple_syrup/domain/conditioning_schedule_selection.py", "simple_syrup/domain/regional_attention_selection.py", diff --git a/tests/test_regional_strategy_comparison_history.py b/tests/regional_generation/regional/test_regional_strategy_comparison_history.py similarity index 100% rename from tests/test_regional_strategy_comparison_history.py rename to tests/regional_generation/regional/test_regional_strategy_comparison_history.py diff --git a/tests/test_regional_strategy_comparison_input_artifacts.py b/tests/regional_generation/regional/test_regional_strategy_comparison_input_artifacts.py similarity index 100% rename from tests/test_regional_strategy_comparison_input_artifacts.py rename to tests/regional_generation/regional/test_regional_strategy_comparison_input_artifacts.py diff --git a/tests/test_regional_strategy_comparison_matrix.py b/tests/regional_generation/regional/test_regional_strategy_comparison_matrix.py similarity index 100% rename from tests/test_regional_strategy_comparison_matrix.py rename to tests/regional_generation/regional/test_regional_strategy_comparison_matrix.py diff --git a/tests/test_regional_strategy_comparison_results.py b/tests/regional_generation/regional/test_regional_strategy_comparison_results.py similarity index 100% rename from tests/test_regional_strategy_comparison_results.py rename to tests/regional_generation/regional/test_regional_strategy_comparison_results.py diff --git a/tests/test_regional_strategy_comparison_validation.py b/tests/regional_generation/regional/test_regional_strategy_comparison_validation.py similarity index 100% rename from tests/test_regional_strategy_comparison_validation.py rename to tests/regional_generation/regional/test_regional_strategy_comparison_validation.py diff --git a/tests/test_regional_strategy_comparison_workflow.py b/tests/regional_generation/regional/test_regional_strategy_comparison_workflow.py similarity index 100% rename from tests/test_regional_strategy_comparison_workflow.py rename to tests/regional_generation/regional/test_regional_strategy_comparison_workflow.py diff --git a/tests/test_regional_tiled_diffusion.py b/tests/regional_generation/regional/test_regional_tiled_diffusion.py similarity index 100% rename from tests/test_regional_tiled_diffusion.py rename to tests/regional_generation/regional/test_regional_tiled_diffusion.py diff --git a/tests/test_regional_tiled_diffusion_sampling_service.py b/tests/regional_generation/regional/test_regional_tiled_diffusion_sampling_service.py similarity index 100% rename from tests/test_regional_tiled_diffusion_sampling_service.py rename to tests/regional_generation/regional/test_regional_tiled_diffusion_sampling_service.py diff --git a/tests/test_regional_visual_benchmark_blind_scoring.py b/tests/regional_generation/regional/test_regional_visual_benchmark_blind_scoring.py similarity index 100% rename from tests/test_regional_visual_benchmark_blind_scoring.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_blind_scoring.py diff --git a/tests/test_regional_visual_benchmark_corpus.py b/tests/regional_generation/regional/test_regional_visual_benchmark_corpus.py similarity index 100% rename from tests/test_regional_visual_benchmark_corpus.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_corpus.py diff --git a/tests/test_regional_visual_benchmark_execution.py b/tests/regional_generation/regional/test_regional_visual_benchmark_execution.py similarity index 100% rename from tests/test_regional_visual_benchmark_execution.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_execution.py diff --git a/tests/test_regional_visual_benchmark_input_artifacts.py b/tests/regional_generation/regional/test_regional_visual_benchmark_input_artifacts.py similarity index 100% rename from tests/test_regional_visual_benchmark_input_artifacts.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_input_artifacts.py diff --git a/tests/test_regional_visual_benchmark_matrix.py b/tests/regional_generation/regional/test_regional_visual_benchmark_matrix.py similarity index 100% rename from tests/test_regional_visual_benchmark_matrix.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_matrix.py diff --git a/tests/test_regional_visual_benchmark_run_artifacts.py b/tests/regional_generation/regional/test_regional_visual_benchmark_run_artifacts.py similarity index 100% rename from tests/test_regional_visual_benchmark_run_artifacts.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_run_artifacts.py diff --git a/tests/test_regional_visual_benchmark_runner.py b/tests/regional_generation/regional/test_regional_visual_benchmark_runner.py similarity index 100% rename from tests/test_regional_visual_benchmark_runner.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_runner.py diff --git a/tests/test_regional_visual_benchmark_workflow.py b/tests/regional_generation/regional/test_regional_visual_benchmark_workflow.py similarity index 100% rename from tests/test_regional_visual_benchmark_workflow.py rename to tests/regional_generation/regional/test_regional_visual_benchmark_workflow.py diff --git a/tests/test_resolved_regional_lora.py b/tests/regional_generation/regional/test_resolved_regional_lora.py similarity index 98% rename from tests/test_resolved_regional_lora.py rename to tests/regional_generation/regional/test_resolved_regional_lora.py index 85637a6..8eb7914 100644 --- a/tests/test_resolved_regional_lora.py +++ b/tests/regional_generation/regional/test_resolved_regional_lora.py @@ -9,10 +9,10 @@ from __future__ import annotations import ast from collections.abc import Callable from dataclasses import FrozenInstanceError, replace -from pathlib import Path from typing import Any, cast import pytest +from support.repository import REPOSITORY_ROOT from simple_syrup.domain.regional_lora_plan import ( RegionalLoraAdapterIdentity, @@ -220,12 +220,7 @@ def test_domain_values_are_frozen_and_import_no_runtime_frameworks() -> None: with pytest.raises(FrozenInstanceError): operation.rank = 8 # type: ignore[misc] - source = ( - Path(__file__).parents[1] - / "simple_syrup" - / "domain" - / "resolved_regional_lora.py" - ) + source = REPOSITORY_ROOT / "simple_syrup" / "domain" / "resolved_regional_lora.py" tree = ast.parse(source.read_text(encoding="utf-8")) imported_roots = { alias.name.split(".", 1)[0] diff --git a/tests/test_resolved_regional_lora_translation.py b/tests/regional_generation/regional/test_resolved_regional_lora_translation.py similarity index 100% rename from tests/test_resolved_regional_lora_translation.py rename to tests/regional_generation/regional/test_resolved_regional_lora_translation.py diff --git a/tests/regional_generation/sdxl/__init__.py b/tests/regional_generation/sdxl/__init__.py new file mode 100644 index 0000000..c95946b --- /dev/null +++ b/tests/regional_generation/sdxl/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation sdxl test behavior.""" diff --git a/tests/regional_generation/sdxl/support/__init__.py b/tests/regional_generation/sdxl/support/__init__.py new file mode 100644 index 0000000..fc0cac9 --- /dev/null +++ b/tests/regional_generation/sdxl/support/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation sdxl support test behavior.""" diff --git a/tests/sdxl_visual_test_inventory.py b/tests/regional_generation/sdxl/support/sdxl_visual_test_inventory.py similarity index 100% rename from tests/sdxl_visual_test_inventory.py rename to tests/regional_generation/sdxl/support/sdxl_visual_test_inventory.py diff --git a/tests/test_sdxl_attention_couple_parity_cases.py b/tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_cases.py similarity index 100% rename from tests/test_sdxl_attention_couple_parity_cases.py rename to tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_cases.py diff --git a/tests/test_sdxl_attention_couple_parity_ownership.py b/tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_ownership.py similarity index 100% rename from tests/test_sdxl_attention_couple_parity_ownership.py rename to tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_ownership.py diff --git a/tests/test_sdxl_attention_couple_parity_results.py b/tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_results.py similarity index 100% rename from tests/test_sdxl_attention_couple_parity_results.py rename to tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_results.py diff --git a/tests/test_sdxl_attention_couple_parity_workflow.py b/tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_workflow.py similarity index 98% rename from tests/test_sdxl_attention_couple_parity_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_workflow.py index 8476d28..f05828c 100644 --- a/tests/test_sdxl_attention_couple_parity_workflow.py +++ b/tests/regional_generation/sdxl/test_sdxl_attention_couple_parity_workflow.py @@ -8,8 +8,6 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.comfy_api import JsonObject from tools.sdxl_attention_couple_parity.cases import ( SdxlAttentionCoupleParityCase, @@ -22,6 +20,10 @@ from tools.sdxl_attention_couple_parity.workflow import ( build_parity_workflow, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_parity_graphs_share_native_sdxl_controls_and_only_change_backend( tmp_path: Path, diff --git a/tests/test_sdxl_attention_coupling_results.py b/tests/regional_generation/sdxl/test_sdxl_attention_coupling_results.py similarity index 100% rename from tests/test_sdxl_attention_coupling_results.py rename to tests/regional_generation/sdxl/test_sdxl_attention_coupling_results.py diff --git a/tests/test_sdxl_attention_coupling_workflow.py b/tests/regional_generation/sdxl/test_sdxl_attention_coupling_workflow.py similarity index 100% rename from tests/test_sdxl_attention_coupling_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_attention_coupling_workflow.py diff --git a/tests/test_sdxl_cold_path_priming.py b/tests/regional_generation/sdxl/test_sdxl_cold_path_priming.py similarity index 100% rename from tests/test_sdxl_cold_path_priming.py rename to tests/regional_generation/sdxl/test_sdxl_cold_path_priming.py diff --git a/tests/test_sdxl_cold_path_results.py b/tests/regional_generation/sdxl/test_sdxl_cold_path_results.py similarity index 100% rename from tests/test_sdxl_cold_path_results.py rename to tests/regional_generation/sdxl/test_sdxl_cold_path_results.py diff --git a/tests/test_sdxl_cold_path_workflow.py b/tests/regional_generation/sdxl/test_sdxl_cold_path_workflow.py similarity index 100% rename from tests/test_sdxl_cold_path_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_cold_path_workflow.py diff --git a/tests/test_sdxl_cold_upstream_profile_workflow.py b/tests/regional_generation/sdxl/test_sdxl_cold_upstream_profile_workflow.py similarity index 100% rename from tests/test_sdxl_cold_upstream_profile_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_cold_upstream_profile_workflow.py diff --git a/tests/test_sdxl_cold_upstream_trace_results.py b/tests/regional_generation/sdxl/test_sdxl_cold_upstream_trace_results.py similarity index 100% rename from tests/test_sdxl_cold_upstream_trace_results.py rename to tests/regional_generation/sdxl/test_sdxl_cold_upstream_trace_results.py diff --git a/tests/test_sdxl_composed_character_controls.py b/tests/regional_generation/sdxl/test_sdxl_composed_character_controls.py similarity index 95% rename from tests/test_sdxl_composed_character_controls.py rename to tests/regional_generation/sdxl/test_sdxl_composed_character_controls.py index 3585ed9..ef801bb 100644 --- a/tests/test_sdxl_composed_character_controls.py +++ b/tests/regional_generation/sdxl/test_sdxl_composed_character_controls.py @@ -8,10 +8,12 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_cases import visual_cases +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_left_adapter_control_preserves_exact_simultaneous_prompt_topology( tmp_path: Path, diff --git a/tests/test_sdxl_composed_lora_completion_cases.py b/tests/regional_generation/sdxl/test_sdxl_composed_lora_completion_cases.py similarity index 98% rename from tests/test_sdxl_composed_lora_completion_cases.py rename to tests/regional_generation/sdxl/test_sdxl_composed_lora_completion_cases.py index 67eb886..7db0f03 100644 --- a/tests/test_sdxl_composed_lora_completion_cases.py +++ b/tests/regional_generation/sdxl/test_sdxl_composed_lora_completion_cases.py @@ -8,8 +8,6 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_adapter_selections import ( LEFT_CHARACTER_SELECTION, RIGHT_CHARACTER_SELECTION, @@ -24,6 +22,10 @@ from tools.sdxl_composed_lora_completion.cases import ( composed_lora_completion_cases, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_completion_cases_lock_causal_order_scope_and_strength(tmp_path: Path) -> None: """Require exact adapter placement for all three unresolved gaps.""" diff --git a/tests/test_sdxl_conventional_variant_workflow.py b/tests/regional_generation/sdxl/test_sdxl_conventional_variant_workflow.py similarity index 100% rename from tests/test_sdxl_conventional_variant_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_conventional_variant_workflow.py diff --git a/tests/test_sdxl_evidence_validation.py b/tests/regional_generation/sdxl/test_sdxl_evidence_validation.py similarity index 100% rename from tests/test_sdxl_evidence_validation.py rename to tests/regional_generation/sdxl/test_sdxl_evidence_validation.py diff --git a/tests/test_sdxl_full_strength_composition_cases.py b/tests/regional_generation/sdxl/test_sdxl_full_strength_composition_cases.py similarity index 96% rename from tests/test_sdxl_full_strength_composition_cases.py rename to tests/regional_generation/sdxl/test_sdxl_full_strength_composition_cases.py index 67bcc9e..98ee464 100644 --- a/tests/test_sdxl_full_strength_composition_cases.py +++ b/tests/regional_generation/sdxl/test_sdxl_full_strength_composition_cases.py @@ -8,8 +8,6 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_adapter_selections import ( LEFT_CHARACTER_SELECTION, RIGHT_CHARACTER_SELECTION, @@ -21,6 +19,10 @@ from tools.sdxl_full_strength_lora_fidelity.composition_cases import ( full_strength_composition_cases, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_full_strength_composition_cases_are_constant_and_ordered( tmp_path: Path, diff --git a/tests/test_sdxl_full_strength_fidelity_cases.py b/tests/regional_generation/sdxl/test_sdxl_full_strength_fidelity_cases.py similarity index 100% rename from tests/test_sdxl_full_strength_fidelity_cases.py rename to tests/regional_generation/sdxl/test_sdxl_full_strength_fidelity_cases.py diff --git a/tests/test_sdxl_full_strength_fidelity_mask.py b/tests/regional_generation/sdxl/test_sdxl_full_strength_fidelity_mask.py similarity index 100% rename from tests/test_sdxl_full_strength_fidelity_mask.py rename to tests/regional_generation/sdxl/test_sdxl_full_strength_fidelity_mask.py diff --git a/tests/test_sdxl_full_strength_fidelity_workflow.py b/tests/regional_generation/sdxl/test_sdxl_full_strength_fidelity_workflow.py similarity index 100% rename from tests/test_sdxl_full_strength_fidelity_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_full_strength_fidelity_workflow.py diff --git a/tests/test_sdxl_global_lora_reference_results.py b/tests/regional_generation/sdxl/test_sdxl_global_lora_reference_results.py similarity index 100% rename from tests/test_sdxl_global_lora_reference_results.py rename to tests/regional_generation/sdxl/test_sdxl_global_lora_reference_results.py diff --git a/tests/test_sdxl_global_lora_reference_workflow.py b/tests/regional_generation/sdxl/test_sdxl_global_lora_reference_workflow.py similarity index 100% rename from tests/test_sdxl_global_lora_reference_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_global_lora_reference_workflow.py diff --git a/tests/test_sdxl_indexed_profile_workflow.py b/tests/regional_generation/sdxl/test_sdxl_indexed_profile_workflow.py similarity index 100% rename from tests/test_sdxl_indexed_profile_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_indexed_profile_workflow.py diff --git a/tests/test_sdxl_integration_checkpoint_link.py b/tests/regional_generation/sdxl/test_sdxl_integration_checkpoint_link.py similarity index 100% rename from tests/test_sdxl_integration_checkpoint_link.py rename to tests/regional_generation/sdxl/test_sdxl_integration_checkpoint_link.py diff --git a/tests/test_sdxl_integration_masks.py b/tests/regional_generation/sdxl/test_sdxl_integration_masks.py similarity index 100% rename from tests/test_sdxl_integration_masks.py rename to tests/regional_generation/sdxl/test_sdxl_integration_masks.py diff --git a/tests/test_sdxl_matched_floor_results.py b/tests/regional_generation/sdxl/test_sdxl_matched_floor_results.py similarity index 100% rename from tests/test_sdxl_matched_floor_results.py rename to tests/regional_generation/sdxl/test_sdxl_matched_floor_results.py diff --git a/tests/test_sdxl_materialization_parity_results.py b/tests/regional_generation/sdxl/test_sdxl_materialization_parity_results.py similarity index 100% rename from tests/test_sdxl_materialization_parity_results.py rename to tests/regional_generation/sdxl/test_sdxl_materialization_parity_results.py diff --git a/tests/test_sdxl_materialization_parity_workflow.py b/tests/regional_generation/sdxl/test_sdxl_materialization_parity_workflow.py similarity index 100% rename from tests/test_sdxl_materialization_parity_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_materialization_parity_workflow.py diff --git a/tests/test_sdxl_multiple_left_cases.py b/tests/regional_generation/sdxl/test_sdxl_multiple_left_cases.py similarity index 99% rename from tests/test_sdxl_multiple_left_cases.py rename to tests/regional_generation/sdxl/test_sdxl_multiple_left_cases.py index 1342f9c..096b2cf 100644 --- a/tests/test_sdxl_multiple_left_cases.py +++ b/tests/regional_generation/sdxl/test_sdxl_multiple_left_cases.py @@ -8,8 +8,6 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_adapter_selections import ( LEFT_CHARACTER_SELECTION, STYLE_SELECTION, @@ -17,6 +15,10 @@ from tools.sdxl_attention_coupling_integration.visual_adapter_selections import from tools.sdxl_attention_coupling_integration.visual_case_model import SdxlVisualCase from tools.sdxl_attention_coupling_integration.visual_cases import visual_cases +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_multiple_left_acceptance_pair_uses_parity_prompt_weight( tmp_path: Path, diff --git a/tests/test_sdxl_native_conditioning_contract.py b/tests/regional_generation/sdxl/test_sdxl_native_conditioning_contract.py similarity index 100% rename from tests/test_sdxl_native_conditioning_contract.py rename to tests/regional_generation/sdxl/test_sdxl_native_conditioning_contract.py diff --git a/tests/test_sdxl_no_lora_comparison_runner.py b/tests/regional_generation/sdxl/test_sdxl_no_lora_comparison_runner.py similarity index 99% rename from tests/test_sdxl_no_lora_comparison_runner.py rename to tests/regional_generation/sdxl/test_sdxl_no_lora_comparison_runner.py index bb697bb..88e669b 100644 --- a/tests/test_sdxl_no_lora_comparison_runner.py +++ b/tests/regional_generation/sdxl/test_sdxl_no_lora_comparison_runner.py @@ -12,7 +12,6 @@ from types import SimpleNamespace from typing import Protocol, cast from pytest import MonkeyPatch -from sdxl_visual_test_inventory import visual_inventory import tools.sdxl_no_lora_comparison.runner as runner from tools.comfy_api import JsonObject @@ -21,6 +20,10 @@ from tools.sdxl_attention_coupling_integration.visual_lora_baseline_cases import SdxlVisualPromptSet, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + class _FakeClient: """Record submitted graphs without invoking ComfyUI.""" diff --git a/tests/test_sdxl_no_lora_comparison_workflow.py b/tests/regional_generation/sdxl/test_sdxl_no_lora_comparison_workflow.py similarity index 100% rename from tests/test_sdxl_no_lora_comparison_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_no_lora_comparison_workflow.py diff --git a/tests/test_sdxl_plain_ksampler_workflow.py b/tests/regional_generation/sdxl/test_sdxl_plain_ksampler_workflow.py similarity index 100% rename from tests/test_sdxl_plain_ksampler_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_plain_ksampler_workflow.py diff --git a/tests/test_sdxl_post_optimization_visual_cases.py b/tests/regional_generation/sdxl/test_sdxl_post_optimization_visual_cases.py similarity index 97% rename from tests/test_sdxl_post_optimization_visual_cases.py rename to tests/regional_generation/sdxl/test_sdxl_post_optimization_visual_cases.py index 33319c1..4aeedbc 100644 --- a/tests/test_sdxl_post_optimization_visual_cases.py +++ b/tests/regional_generation/sdxl/test_sdxl_post_optimization_visual_cases.py @@ -9,8 +9,6 @@ from __future__ import annotations from dataclasses import replace from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_adapter_selections import ( LEFT_CHARACTER_SELECTION, RIGHT_CHARACTER_SELECTION, @@ -24,6 +22,10 @@ from tools.sdxl_post_optimization_visual_proof.cases import ( ) from tools.sdxl_two_character_oracle import two_character_oracle_cases +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_post_optimization_cases_lock_placement_strength_and_style_prefix( tmp_path: Path, diff --git a/tests/test_sdxl_regional_lora_performance_measurement.py b/tests/regional_generation/sdxl/test_sdxl_regional_lora_performance_measurement.py similarity index 100% rename from tests/test_sdxl_regional_lora_performance_measurement.py rename to tests/regional_generation/sdxl/test_sdxl_regional_lora_performance_measurement.py diff --git a/tests/test_sdxl_regional_lora_performance_workflow.py b/tests/regional_generation/sdxl/test_sdxl_regional_lora_performance_workflow.py similarity index 100% rename from tests/test_sdxl_regional_lora_performance_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_regional_lora_performance_workflow.py diff --git a/tests/test_sdxl_regional_negative_parity_workflow.py b/tests/regional_generation/sdxl/test_sdxl_regional_negative_parity_workflow.py similarity index 100% rename from tests/test_sdxl_regional_negative_parity_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_regional_negative_parity_workflow.py diff --git a/tests/test_sdxl_regional_scaling_cases.py b/tests/regional_generation/sdxl/test_sdxl_regional_scaling_cases.py similarity index 100% rename from tests/test_sdxl_regional_scaling_cases.py rename to tests/regional_generation/sdxl/test_sdxl_regional_scaling_cases.py diff --git a/tests/test_sdxl_regional_scaling_profile_runner.py b/tests/regional_generation/sdxl/test_sdxl_regional_scaling_profile_runner.py similarity index 100% rename from tests/test_sdxl_regional_scaling_profile_runner.py rename to tests/regional_generation/sdxl/test_sdxl_regional_scaling_profile_runner.py diff --git a/tests/test_sdxl_regional_scaling_results.py b/tests/regional_generation/sdxl/test_sdxl_regional_scaling_results.py similarity index 100% rename from tests/test_sdxl_regional_scaling_results.py rename to tests/regional_generation/sdxl/test_sdxl_regional_scaling_results.py diff --git a/tests/test_sdxl_regional_scaling_workflow.py b/tests/regional_generation/sdxl/test_sdxl_regional_scaling_workflow.py similarity index 100% rename from tests/test_sdxl_regional_scaling_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_regional_scaling_workflow.py diff --git a/tests/test_sdxl_schedule_mask_completion_cases.py b/tests/regional_generation/sdxl/test_sdxl_schedule_mask_completion_cases.py similarity index 97% rename from tests/test_sdxl_schedule_mask_completion_cases.py rename to tests/regional_generation/sdxl/test_sdxl_schedule_mask_completion_cases.py index 4b02cb6..5352e17 100644 --- a/tests/test_sdxl_schedule_mask_completion_cases.py +++ b/tests/regional_generation/sdxl/test_sdxl_schedule_mask_completion_cases.py @@ -8,8 +8,6 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_adapter_selections import ( LEFT_CHARACTER_SELECTION, RIGHT_CHARACTER_SELECTION, @@ -27,6 +25,10 @@ from tools.sdxl_schedule_mask_completion.cases import ( schedule_mask_completion_cases, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_cases_isolate_independent_schedules_and_mask_geometry(tmp_path: Path) -> None: """Keep one sampling axis different in each focused RA-06 case.""" diff --git a/tests/test_sdxl_spatial_mode_oracle_case.py b/tests/regional_generation/sdxl/test_sdxl_spatial_mode_oracle_case.py similarity index 95% rename from tests/test_sdxl_spatial_mode_oracle_case.py rename to tests/regional_generation/sdxl/test_sdxl_spatial_mode_oracle_case.py index 56d7a2d..55ee6fe 100644 --- a/tests/test_sdxl_spatial_mode_oracle_case.py +++ b/tests/regional_generation/sdxl/test_sdxl_spatial_mode_oracle_case.py @@ -9,8 +9,6 @@ from __future__ import annotations from dataclasses import replace from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_case_model import VisualMode from tools.sdxl_attention_coupling_integration.visual_lora_baseline_cases import ( SdxlVisualPromptSet, @@ -21,6 +19,10 @@ from tools.sdxl_spatial_mode_oracle import ( ) from tools.sdxl_two_character_oracle import two_character_oracle_cases +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_spatial_oracle_changes_only_identity_label_and_declared_modes( tmp_path: Path, diff --git a/tests/test_sdxl_steady_state_measurement.py b/tests/regional_generation/sdxl/test_sdxl_steady_state_measurement.py similarity index 100% rename from tests/test_sdxl_steady_state_measurement.py rename to tests/regional_generation/sdxl/test_sdxl_steady_state_measurement.py diff --git a/tests/test_sdxl_steady_state_runner.py b/tests/regional_generation/sdxl/test_sdxl_steady_state_runner.py similarity index 100% rename from tests/test_sdxl_steady_state_runner.py rename to tests/regional_generation/sdxl/test_sdxl_steady_state_runner.py diff --git a/tests/test_sdxl_two_adapter_performance_workflow.py b/tests/regional_generation/sdxl/test_sdxl_two_adapter_performance_workflow.py similarity index 100% rename from tests/test_sdxl_two_adapter_performance_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_two_adapter_performance_workflow.py diff --git a/tests/test_sdxl_two_character_oracle_case.py b/tests/regional_generation/sdxl/test_sdxl_two_character_oracle_case.py similarity index 97% rename from tests/test_sdxl_two_character_oracle_case.py rename to tests/regional_generation/sdxl/test_sdxl_two_character_oracle_case.py index 408b05e..199cfc5 100644 --- a/tests/test_sdxl_two_character_oracle_case.py +++ b/tests/regional_generation/sdxl/test_sdxl_two_character_oracle_case.py @@ -8,8 +8,6 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.sampling_controls import ( SDXL_VISUAL_SAMPLING, ) @@ -29,6 +27,10 @@ from tools.sdxl_two_character_oracle import ( two_character_oracle_cases, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_two_character_oracle_locks_accepted_graph_without_model_identity( tmp_path: Path, diff --git a/tests/test_sdxl_visual_artifacts.py b/tests/regional_generation/sdxl/test_sdxl_visual_artifacts.py similarity index 100% rename from tests/test_sdxl_visual_artifacts.py rename to tests/regional_generation/sdxl/test_sdxl_visual_artifacts.py diff --git a/tests/test_sdxl_visual_cases.py b/tests/regional_generation/sdxl/test_sdxl_visual_cases.py similarity index 99% rename from tests/test_sdxl_visual_cases.py rename to tests/regional_generation/sdxl/test_sdxl_visual_cases.py index 8436ffb..932bc81 100644 --- a/tests/test_sdxl_visual_cases.py +++ b/tests/regional_generation/sdxl/test_sdxl_visual_cases.py @@ -8,8 +8,6 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.sdxl_attention_coupling_integration.visual_adapter_selections import ( RIGHT_CHARACTER_SELECTION, STYLE_SELECTION, @@ -29,6 +27,10 @@ from tools.sdxl_attention_coupling_integration.visual_prompt_defaults import ( RIGHT_BASE_L, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_visual_baseline_uses_explicit_multi_subject_tags(tmp_path: Path) -> None: """Keep the ownership matrix anchored to a visible two-subject composition.""" diff --git a/tests/test_sdxl_visual_comfy_model_root.py b/tests/regional_generation/sdxl/test_sdxl_visual_comfy_model_root.py similarity index 100% rename from tests/test_sdxl_visual_comfy_model_root.py rename to tests/regional_generation/sdxl/test_sdxl_visual_comfy_model_root.py diff --git a/tests/test_sdxl_visual_conditioning.py b/tests/regional_generation/sdxl/test_sdxl_visual_conditioning.py similarity index 99% rename from tests/test_sdxl_visual_conditioning.py rename to tests/regional_generation/sdxl/test_sdxl_visual_conditioning.py index cb8d09c..d54e581 100644 --- a/tests/test_sdxl_visual_conditioning.py +++ b/tests/regional_generation/sdxl/test_sdxl_visual_conditioning.py @@ -10,8 +10,6 @@ import json from pathlib import Path from typing import cast -from sdxl_visual_test_inventory import visual_inventory - from tools.comfy_api import JsonObject from tools.sdxl_attention_coupling_integration.graph import SdxlWorkflowGraph from tools.sdxl_attention_coupling_integration.visual_cases import visual_cases @@ -24,6 +22,10 @@ from tools.sdxl_attention_coupling_integration.visual_prompt_defaults import ( BASE_POSITIVE_L, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_global_style_and_regional_character_use_distinct_comfy_paths( tmp_path: Path, diff --git a/tests/test_sdxl_visual_diagnostic_history.py b/tests/regional_generation/sdxl/test_sdxl_visual_diagnostic_history.py similarity index 100% rename from tests/test_sdxl_visual_diagnostic_history.py rename to tests/regional_generation/sdxl/test_sdxl_visual_diagnostic_history.py diff --git a/tests/test_sdxl_visual_diagnostic_validation.py b/tests/regional_generation/sdxl/test_sdxl_visual_diagnostic_validation.py similarity index 100% rename from tests/test_sdxl_visual_diagnostic_validation.py rename to tests/regional_generation/sdxl/test_sdxl_visual_diagnostic_validation.py diff --git a/tests/test_sdxl_visual_diagnostic_workflow.py b/tests/regional_generation/sdxl/test_sdxl_visual_diagnostic_workflow.py similarity index 100% rename from tests/test_sdxl_visual_diagnostic_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_visual_diagnostic_workflow.py diff --git a/tests/test_sdxl_visual_inventory.py b/tests/regional_generation/sdxl/test_sdxl_visual_inventory.py similarity index 97% rename from tests/test_sdxl_visual_inventory.py rename to tests/regional_generation/sdxl/test_sdxl_visual_inventory.py index 6e64b79..8bdb1a3 100644 --- a/tests/test_sdxl_visual_inventory.py +++ b/tests/regional_generation/sdxl/test_sdxl_visual_inventory.py @@ -10,12 +10,15 @@ import json from pathlib import Path import pytest -from sdxl_visual_test_inventory import visual_inventory from tools.sdxl_attention_coupling_integration.visual_inventory import ( SdxlVisualInventory, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_inventory_loads_complete_anonymous_roles(tmp_path: Path) -> None: """Load external artifact paths, labels, and triggers without fixed identities.""" diff --git a/tests/test_sdxl_visual_launch.py b/tests/regional_generation/sdxl/test_sdxl_visual_launch.py similarity index 100% rename from tests/test_sdxl_visual_launch.py rename to tests/regional_generation/sdxl/test_sdxl_visual_launch.py diff --git a/tests/test_sdxl_visual_lora_baseline_cases.py b/tests/regional_generation/sdxl/test_sdxl_visual_lora_baseline_cases.py similarity index 98% rename from tests/test_sdxl_visual_lora_baseline_cases.py rename to tests/regional_generation/sdxl/test_sdxl_visual_lora_baseline_cases.py index a56dee4..599ab15 100644 --- a/tests/test_sdxl_visual_lora_baseline_cases.py +++ b/tests/regional_generation/sdxl/test_sdxl_visual_lora_baseline_cases.py @@ -9,8 +9,6 @@ from __future__ import annotations from pathlib import Path from typing import cast -from sdxl_visual_test_inventory import visual_inventory - from tools.comfy_api import JsonObject from tools.sdxl_attention_coupling_integration.graph import SdxlWorkflowGraph from tools.sdxl_attention_coupling_integration.visual_conditioning import ( @@ -21,6 +19,10 @@ from tools.sdxl_attention_coupling_integration.visual_lora_baseline_cases import lora_baseline_cases, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_baseline_cases_swap_only_the_active_character_sections( tmp_path: Path, diff --git a/tests/test_sdxl_visual_lora_graph.py b/tests/regional_generation/sdxl/test_sdxl_visual_lora_graph.py similarity index 100% rename from tests/test_sdxl_visual_lora_graph.py rename to tests/regional_generation/sdxl/test_sdxl_visual_lora_graph.py diff --git a/tests/test_sdxl_visual_masks.py b/tests/regional_generation/sdxl/test_sdxl_visual_masks.py similarity index 100% rename from tests/test_sdxl_visual_masks.py rename to tests/regional_generation/sdxl/test_sdxl_visual_masks.py diff --git a/tests/test_sdxl_visual_model_links.py b/tests/regional_generation/sdxl/test_sdxl_visual_model_links.py similarity index 100% rename from tests/test_sdxl_visual_model_links.py rename to tests/regional_generation/sdxl/test_sdxl_visual_model_links.py diff --git a/tests/test_sdxl_visual_prompt_fixture.py b/tests/regional_generation/sdxl/test_sdxl_visual_prompt_fixture.py similarity index 100% rename from tests/test_sdxl_visual_prompt_fixture.py rename to tests/regional_generation/sdxl/test_sdxl_visual_prompt_fixture.py diff --git a/tests/test_sdxl_visual_registration_probe.py b/tests/regional_generation/sdxl/test_sdxl_visual_registration_probe.py similarity index 95% rename from tests/test_sdxl_visual_registration_probe.py rename to tests/regional_generation/sdxl/test_sdxl_visual_registration_probe.py index e8ada6f..90c5f79 100644 --- a/tests/test_sdxl_visual_registration_probe.py +++ b/tests/regional_generation/sdxl/test_sdxl_visual_registration_probe.py @@ -9,7 +9,6 @@ from __future__ import annotations from pathlib import Path import pytest -from sdxl_visual_test_inventory import visual_inventory from tools.comfy_api import JsonObject from tools.prove_sdxl_visual_registration import ( @@ -21,6 +20,10 @@ from tools.sdxl_attention_coupling_integration.visual_adapter_selections import CHECKPOINT_SELECTION, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_required_nodes_cover_all_public_sampler_geometries(tmp_path: Path) -> None: """Probe the same live registration surface as the 14-output matrix.""" diff --git a/tests/test_sdxl_visual_results.py b/tests/regional_generation/sdxl/test_sdxl_visual_results.py similarity index 99% rename from tests/test_sdxl_visual_results.py rename to tests/regional_generation/sdxl/test_sdxl_visual_results.py index 71ad071..1a39268 100644 --- a/tests/test_sdxl_visual_results.py +++ b/tests/regional_generation/sdxl/test_sdxl_visual_results.py @@ -13,7 +13,6 @@ from pathlib import Path from typing import cast from PIL import Image -from sdxl_visual_test_inventory import visual_inventory from tools.comfy_api import JsonObject from tools.sdxl_attention_coupling_integration.matrix import MODES, SdxlIntegrationMode @@ -41,6 +40,10 @@ from tools.sdxl_attention_coupling_integration.visual_workflow import ( build_sdxl_visual_workflow, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_recorder_serializes_prompts_from_the_exact_case(tmp_path: Path) -> None: """Keep external global and regional negatives in durable result evidence.""" diff --git a/tests/test_sdxl_visual_runtime_expectations.py b/tests/regional_generation/sdxl/test_sdxl_visual_runtime_expectations.py similarity index 100% rename from tests/test_sdxl_visual_runtime_expectations.py rename to tests/regional_generation/sdxl/test_sdxl_visual_runtime_expectations.py diff --git a/tests/test_sdxl_visual_sampling_controls.py b/tests/regional_generation/sdxl/test_sdxl_visual_sampling_controls.py similarity index 100% rename from tests/test_sdxl_visual_sampling_controls.py rename to tests/regional_generation/sdxl/test_sdxl_visual_sampling_controls.py diff --git a/tests/test_sdxl_visual_workflow.py b/tests/regional_generation/sdxl/test_sdxl_visual_workflow.py similarity index 98% rename from tests/test_sdxl_visual_workflow.py rename to tests/regional_generation/sdxl/test_sdxl_visual_workflow.py index 43d38ee..f87b26f 100644 --- a/tests/test_sdxl_visual_workflow.py +++ b/tests/regional_generation/sdxl/test_sdxl_visual_workflow.py @@ -8,14 +8,16 @@ from __future__ import annotations from pathlib import Path -from sdxl_visual_test_inventory import visual_inventory - from tools.comfy_api import JsonObject from tools.sdxl_attention_coupling_integration.visual_cases import visual_cases from tools.sdxl_attention_coupling_integration.visual_workflow import ( build_sdxl_visual_workflow, ) +from .support.sdxl_visual_test_inventory import ( + visual_inventory, +) + def test_baseline_builds_only_one_native_full_sampler(tmp_path: Path) -> None: """Keep one source trajectory and four authored G/L encodes.""" diff --git a/tests/regional_generation/spatial/__init__.py b/tests/regional_generation/spatial/__init__.py new file mode 100644 index 0000000..ec81653 --- /dev/null +++ b/tests/regional_generation/spatial/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation spatial test behavior.""" diff --git a/tests/test_contextual_diffusion_domain.py b/tests/regional_generation/spatial/test_contextual_diffusion_domain.py similarity index 100% rename from tests/test_contextual_diffusion_domain.py rename to tests/regional_generation/spatial/test_contextual_diffusion_domain.py diff --git a/tests/test_contextual_diffusion_sampling_boundary.py b/tests/regional_generation/spatial/test_contextual_diffusion_sampling_boundary.py similarity index 100% rename from tests/test_contextual_diffusion_sampling_boundary.py rename to tests/regional_generation/spatial/test_contextual_diffusion_sampling_boundary.py diff --git a/tests/test_contextual_diffusion_sampling_service.py b/tests/regional_generation/spatial/test_contextual_diffusion_sampling_service.py similarity index 100% rename from tests/test_contextual_diffusion_sampling_service.py rename to tests/regional_generation/spatial/test_contextual_diffusion_sampling_service.py diff --git a/tests/test_contextual_integration_diagnostics.py b/tests/regional_generation/spatial/test_contextual_integration_diagnostics.py similarity index 100% rename from tests/test_contextual_integration_diagnostics.py rename to tests/regional_generation/spatial/test_contextual_integration_diagnostics.py diff --git a/tests/test_contextual_model_wrapper.py b/tests/regional_generation/spatial/test_contextual_model_wrapper.py similarity index 100% rename from tests/test_contextual_model_wrapper.py rename to tests/regional_generation/spatial/test_contextual_model_wrapper.py diff --git a/tests/test_spatial_model_arguments_validation.py b/tests/regional_generation/spatial/test_spatial_model_arguments_validation.py similarity index 100% rename from tests/test_spatial_model_arguments_validation.py rename to tests/regional_generation/spatial/test_spatial_model_arguments_validation.py diff --git a/tests/test_spatial_tensor_projection_characterization.py b/tests/regional_generation/spatial/test_spatial_tensor_projection_characterization.py similarity index 100% rename from tests/test_spatial_tensor_projection_characterization.py rename to tests/regional_generation/spatial/test_spatial_tensor_projection_characterization.py diff --git a/tests/test_spatial_view_model_arguments.py b/tests/regional_generation/spatial/test_spatial_view_model_arguments.py similarity index 100% rename from tests/test_spatial_view_model_arguments.py rename to tests/regional_generation/spatial/test_spatial_view_model_arguments.py diff --git a/tests/test_spatial_views.py b/tests/regional_generation/spatial/test_spatial_views.py similarity index 100% rename from tests/test_spatial_views.py rename to tests/regional_generation/spatial/test_spatial_views.py diff --git a/tests/regional_generation/unet/__init__.py b/tests/regional_generation/unet/__init__.py new file mode 100644 index 0000000..32fa083 --- /dev/null +++ b/tests/regional_generation/unet/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation unet test behavior.""" diff --git a/tests/regional_generation/unet/support/__init__.py b/tests/regional_generation/unet/support/__init__.py new file mode 100644 index 0000000..dad0afd --- /dev/null +++ b/tests/regional_generation/unet/support/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own regional generation unet support test behavior.""" diff --git a/tests/unet_attention_coupling_call_harness.py b/tests/regional_generation/unet/support/unet_attention_coupling_call_harness.py similarity index 98% rename from tests/unet_attention_coupling_call_harness.py rename to tests/regional_generation/unet/support/unet_attention_coupling_call_harness.py index b23106d..275a2b8 100644 --- a/tests/unet_attention_coupling_call_harness.py +++ b/tests/regional_generation/unet/support/unet_attention_coupling_call_harness.py @@ -9,11 +9,6 @@ from __future__ import annotations from dataclasses import dataclass import torch -from attention_coupling_invariant_contract import ( - AttentionCouplingCallHarness, - AttentionCouplingCallObservation, - AttentionCouplingCallScenario, -) from comfy.patcher_extension import WrappersMP from simple_syrup.domain.regional_attention_batch import ( @@ -33,6 +28,12 @@ from simple_syrup.runtime.regional_attention_diagnostics import ( RegionalAttentionDiagnosticsBuilder, ) +from ...attention_coupling.support.attention_coupling_invariant_contract import ( + AttentionCouplingCallHarness, + AttentionCouplingCallObservation, + AttentionCouplingCallScenario, +) + class _UnetDiffusionModel(torch.nn.Module): """Provide weak-referenceable UNet model identity for call invariants.""" diff --git a/tests/unet_attention_coupling_diagnostics_harness.py b/tests/regional_generation/unet/support/unet_attention_coupling_diagnostics_harness.py similarity index 79% rename from tests/unet_attention_coupling_diagnostics_harness.py rename to tests/regional_generation/unet/support/unet_attention_coupling_diagnostics_harness.py index 221d7a8..9257994 100644 --- a/tests/unet_attention_coupling_diagnostics_harness.py +++ b/tests/regional_generation/unet/support/unet_attention_coupling_diagnostics_harness.py @@ -6,12 +6,15 @@ from __future__ import annotations -from attention_coupling_diagnostics_harness import build_diagnostics_observation -from attention_coupling_invariant_contract import ( +from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel + +from ...attention_coupling.support.attention_coupling_diagnostics_harness import ( + build_diagnostics_observation, +) +from ...attention_coupling.support.attention_coupling_invariant_contract import ( AttentionCouplingDiagnosticsHarness, AttentionCouplingDiagnosticsObservation, ) -from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel class UnetAttentionCouplingDiagnosticsHarness(AttentionCouplingDiagnosticsHarness): diff --git a/tests/unet_attention_coupling_invariant_harness.py b/tests/regional_generation/unet/support/unet_attention_coupling_invariant_harness.py similarity index 98% rename from tests/unet_attention_coupling_invariant_harness.py rename to tests/regional_generation/unet/support/unet_attention_coupling_invariant_harness.py index 5ceb752..e32ee16 100644 --- a/tests/unet_attention_coupling_invariant_harness.py +++ b/tests/regional_generation/unet/support/unet_attention_coupling_invariant_harness.py @@ -7,12 +7,6 @@ from __future__ import annotations import torch -from attention_coupling_invariant_contract import ( - AttentionCouplingInvariantEntry, - AttentionCouplingInvariantHarness, - AttentionCouplingInvariantObservation, - AttentionCouplingInvariantScenario, -) from simple_syrup.domain.regional_attention import RegionalAttentionBranch from simple_syrup.domain.regional_attention_batch import ( @@ -25,6 +19,13 @@ from simple_syrup.runtime.attention_coupling.unet_attn2_execution import ( UnetAttn2Execution, ) +from ...attention_coupling.support.attention_coupling_invariant_contract import ( + AttentionCouplingInvariantEntry, + AttentionCouplingInvariantHarness, + AttentionCouplingInvariantObservation, + AttentionCouplingInvariantScenario, +) + class _DeterministicUnetAttention: """Evaluate one packed attn2 batch with deterministic context values.""" diff --git a/tests/unet_attention_coupling_lifecycle_harness.py b/tests/regional_generation/unet/support/unet_attention_coupling_lifecycle_harness.py similarity index 94% rename from tests/unet_attention_coupling_lifecycle_harness.py rename to tests/regional_generation/unet/support/unet_attention_coupling_lifecycle_harness.py index 1c2aaca..1dd46fa 100644 --- a/tests/unet_attention_coupling_lifecycle_harness.py +++ b/tests/regional_generation/unet/support/unet_attention_coupling_lifecycle_harness.py @@ -10,11 +10,6 @@ from copy import deepcopy from typing import Any import torch -from attention_coupling_invariant_contract import ( - AttentionCouplingLifecycleHarness, - AttentionCouplingLifecycleObservation, -) -from attention_coupling_invariant_values import scheduled_invariant_plan from torch import nn from simple_syrup.runtime.attention_coupling.unet import ( @@ -31,6 +26,14 @@ from simple_syrup.runtime.regional_lora.standard_unet_operation_preparation impo ) from simple_syrup.runtime.regional_lora_plan_adapter import RegionalLoraPlanAdaptation +from ...attention_coupling.support.attention_coupling_invariant_contract import ( + AttentionCouplingLifecycleHarness, + AttentionCouplingLifecycleObservation, +) +from ...attention_coupling.support.attention_coupling_invariant_values import ( + scheduled_invariant_plan, +) + class _UnetModelRoot(nn.Module): """Expose a diffusion model at Comfy's standard MODEL path.""" diff --git a/tests/test_standard_unet_attention_weighting.py b/tests/regional_generation/unet/test_standard_unet_attention_weighting.py similarity index 100% rename from tests/test_standard_unet_attention_weighting.py rename to tests/regional_generation/unet/test_standard_unet_attention_weighting.py diff --git a/tests/test_standard_unet_cold_diagnostics.py b/tests/regional_generation/unet/test_standard_unet_cold_diagnostics.py similarity index 100% rename from tests/test_standard_unet_cold_diagnostics.py rename to tests/regional_generation/unet/test_standard_unet_cold_diagnostics.py diff --git a/tests/test_standard_unet_cold_sampling.py b/tests/regional_generation/unet/test_standard_unet_cold_sampling.py similarity index 100% rename from tests/test_standard_unet_cold_sampling.py rename to tests/regional_generation/unet/test_standard_unet_cold_sampling.py diff --git a/tests/test_standard_unet_composition_diagnostics.py b/tests/regional_generation/unet/test_standard_unet_composition_diagnostics.py similarity index 100% rename from tests/test_standard_unet_composition_diagnostics.py rename to tests/regional_generation/unet/test_standard_unet_composition_diagnostics.py diff --git a/tests/test_standard_unet_geometry_matrix.py b/tests/regional_generation/unet/test_standard_unet_geometry_matrix.py similarity index 100% rename from tests/test_standard_unet_geometry_matrix.py rename to tests/regional_generation/unet/test_standard_unet_geometry_matrix.py diff --git a/tests/test_standard_unet_independent_schedules.py b/tests/regional_generation/unet/test_standard_unet_independent_schedules.py similarity index 100% rename from tests/test_standard_unet_independent_schedules.py rename to tests/regional_generation/unet/test_standard_unet_independent_schedules.py diff --git a/tests/test_standard_unet_model_capability_detection.py b/tests/regional_generation/unet/test_standard_unet_model_capability_detection.py similarity index 98% rename from tests/test_standard_unet_model_capability_detection.py rename to tests/regional_generation/unet/test_standard_unet_model_capability_detection.py index ad6ffef..8a29bc8 100644 --- a/tests/test_standard_unet_model_capability_detection.py +++ b/tests/regional_generation/unet/test_standard_unet_model_capability_detection.py @@ -13,12 +13,6 @@ import comfy.model_base import pytest import torch from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel -from regional_model_capability_test_values import ( - AlternateImageLatent, - TemporalLatent, - patcher, - standard_unet_graph, -) from simple_syrup.domain.regional_model_capabilities import ( RegionalAttentionBackend, @@ -30,6 +24,13 @@ from simple_syrup.runtime.standard_unet_model_capability import ( StandardUnetModelCapabilityDetector, ) +from ..regional.support.regional_model_capability_test_values import ( + AlternateImageLatent, + TemporalLatent, + patcher, + standard_unet_graph, +) + class _GraphOptions(TypedDict, total=False): """Describe supported graph mutations for rejection characterization.""" diff --git a/tests/test_standard_unet_model_output_validation.py b/tests/regional_generation/unet/test_standard_unet_model_output_validation.py similarity index 100% rename from tests/test_standard_unet_model_output_validation.py rename to tests/regional_generation/unet/test_standard_unet_model_output_validation.py diff --git a/tests/test_standard_unet_native_admission.py b/tests/regional_generation/unet/test_standard_unet_native_admission.py similarity index 100% rename from tests/test_standard_unet_native_admission.py rename to tests/regional_generation/unet/test_standard_unet_native_admission.py diff --git a/tests/test_standard_unet_operation_schedule.py b/tests/regional_generation/unet/test_standard_unet_operation_schedule.py similarity index 100% rename from tests/test_standard_unet_operation_schedule.py rename to tests/regional_generation/unet/test_standard_unet_operation_schedule.py diff --git a/tests/test_standard_unet_packed_operation_masks.py b/tests/regional_generation/unet/test_standard_unet_packed_operation_masks.py similarity index 100% rename from tests/test_standard_unet_packed_operation_masks.py rename to tests/regional_generation/unet/test_standard_unet_packed_operation_masks.py diff --git a/tests/test_standard_unet_reference_attention_oracle.py b/tests/regional_generation/unet/test_standard_unet_reference_attention_oracle.py similarity index 100% rename from tests/test_standard_unet_reference_attention_oracle.py rename to tests/regional_generation/unet/test_standard_unet_reference_attention_oracle.py diff --git a/tests/test_standard_unet_shared_trajectory_integration.py b/tests/regional_generation/unet/test_standard_unet_shared_trajectory_integration.py similarity index 100% rename from tests/test_standard_unet_shared_trajectory_integration.py rename to tests/regional_generation/unet/test_standard_unet_shared_trajectory_integration.py diff --git a/tests/test_standard_unet_target_capabilities.py b/tests/regional_generation/unet/test_standard_unet_target_capabilities.py similarity index 100% rename from tests/test_standard_unet_target_capabilities.py rename to tests/regional_generation/unet/test_standard_unet_target_capabilities.py diff --git a/tests/test_standard_unet_variant_base_attention.py b/tests/regional_generation/unet/test_standard_unet_variant_base_attention.py similarity index 100% rename from tests/test_standard_unet_variant_base_attention.py rename to tests/regional_generation/unet/test_standard_unet_variant_base_attention.py diff --git a/tests/test_standard_unet_variant_conditioning.py b/tests/regional_generation/unet/test_standard_unet_variant_conditioning.py similarity index 100% rename from tests/test_standard_unet_variant_conditioning.py rename to tests/regional_generation/unet/test_standard_unet_variant_conditioning.py diff --git a/tests/test_standard_unet_variant_execution.py b/tests/regional_generation/unet/test_standard_unet_variant_execution.py similarity index 100% rename from tests/test_standard_unet_variant_execution.py rename to tests/regional_generation/unet/test_standard_unet_variant_execution.py diff --git a/tests/test_standard_unet_variant_graph_attention.py b/tests/regional_generation/unet/test_standard_unet_variant_graph_attention.py similarity index 100% rename from tests/test_standard_unet_variant_graph_attention.py rename to tests/regional_generation/unet/test_standard_unet_variant_graph_attention.py diff --git a/tests/test_standard_unet_variant_invocation.py b/tests/regional_generation/unet/test_standard_unet_variant_invocation.py similarity index 100% rename from tests/test_standard_unet_variant_invocation.py rename to tests/regional_generation/unet/test_standard_unet_variant_invocation.py diff --git a/tests/test_standard_unet_variant_lane_graph_attention.py b/tests/regional_generation/unet/test_standard_unet_variant_lane_graph_attention.py similarity index 100% rename from tests/test_standard_unet_variant_lane_graph_attention.py rename to tests/regional_generation/unet/test_standard_unet_variant_lane_graph_attention.py diff --git a/tests/test_standard_unet_variant_lane_plan.py b/tests/regional_generation/unet/test_standard_unet_variant_lane_plan.py similarity index 100% rename from tests/test_standard_unet_variant_lane_plan.py rename to tests/regional_generation/unet/test_standard_unet_variant_lane_plan.py diff --git a/tests/test_standard_unet_variant_materialization.py b/tests/regional_generation/unet/test_standard_unet_variant_materialization.py similarity index 100% rename from tests/test_standard_unet_variant_materialization.py rename to tests/regional_generation/unet/test_standard_unet_variant_materialization.py diff --git a/tests/test_standard_unet_variant_materialization_device.py b/tests/regional_generation/unet/test_standard_unet_variant_materialization_device.py similarity index 100% rename from tests/test_standard_unet_variant_materialization_device.py rename to tests/regional_generation/unet/test_standard_unet_variant_materialization_device.py diff --git a/tests/test_standard_unet_variant_output.py b/tests/regional_generation/unet/test_standard_unet_variant_output.py similarity index 100% rename from tests/test_standard_unet_variant_output.py rename to tests/regional_generation/unet/test_standard_unet_variant_output.py diff --git a/tests/test_standard_unet_variant_residency_handoff.py b/tests/regional_generation/unet/test_standard_unet_variant_residency_handoff.py similarity index 100% rename from tests/test_standard_unet_variant_residency_handoff.py rename to tests/regional_generation/unet/test_standard_unet_variant_residency_handoff.py diff --git a/tests/test_standard_unet_variant_root.py b/tests/regional_generation/unet/test_standard_unet_variant_root.py similarity index 100% rename from tests/test_standard_unet_variant_root.py rename to tests/regional_generation/unet/test_standard_unet_variant_root.py diff --git a/tests/test_standard_unet_variant_shell.py b/tests/regional_generation/unet/test_standard_unet_variant_shell.py similarity index 100% rename from tests/test_standard_unet_variant_shell.py rename to tests/regional_generation/unet/test_standard_unet_variant_shell.py diff --git a/tests/test_standard_unet_variant_spatial_context.py b/tests/regional_generation/unet/test_standard_unet_variant_spatial_context.py similarity index 100% rename from tests/test_standard_unet_variant_spatial_context.py rename to tests/regional_generation/unet/test_standard_unet_variant_spatial_context.py diff --git a/tests/test_standard_unet_variant_static_residency.py b/tests/regional_generation/unet/test_standard_unet_variant_static_residency.py similarity index 100% rename from tests/test_standard_unet_variant_static_residency.py rename to tests/regional_generation/unet/test_standard_unet_variant_static_residency.py diff --git a/tests/test_standard_unet_variant_topology.py b/tests/regional_generation/unet/test_standard_unet_variant_topology.py similarity index 100% rename from tests/test_standard_unet_variant_topology.py rename to tests/regional_generation/unet/test_standard_unet_variant_topology.py diff --git a/tests/test_unet_attention_backend.py b/tests/regional_generation/unet/test_unet_attention_backend.py similarity index 100% rename from tests/test_unet_attention_backend.py rename to tests/regional_generation/unet/test_unet_attention_backend.py diff --git a/tests/test_unet_attention_backend_integration.py b/tests/regional_generation/unet/test_unet_attention_backend_integration.py similarity index 100% rename from tests/test_unet_attention_backend_integration.py rename to tests/regional_generation/unet/test_unet_attention_backend_integration.py diff --git a/tests/test_unet_attention_branch_batch.py b/tests/regional_generation/unet/test_unet_attention_branch_batch.py similarity index 100% rename from tests/test_unet_attention_branch_batch.py rename to tests/regional_generation/unet/test_unet_attention_branch_batch.py diff --git a/tests/test_unet_attention_context_wrapper.py b/tests/regional_generation/unet/test_unet_attention_context_wrapper.py similarity index 100% rename from tests/test_unet_attention_context_wrapper.py rename to tests/regional_generation/unet/test_unet_attention_context_wrapper.py diff --git a/tests/test_unet_attention_coupling_model_family.py b/tests/regional_generation/unet/test_unet_attention_coupling_model_family.py similarity index 100% rename from tests/test_unet_attention_coupling_model_family.py rename to tests/regional_generation/unet/test_unet_attention_coupling_model_family.py diff --git a/tests/test_unet_attention_diagnostics.py b/tests/regional_generation/unet/test_unet_attention_diagnostics.py similarity index 100% rename from tests/test_unet_attention_diagnostics.py rename to tests/regional_generation/unet/test_unet_attention_diagnostics.py diff --git a/tests/test_unet_attention_geometry.py b/tests/regional_generation/unet/test_unet_attention_geometry.py similarity index 100% rename from tests/test_unet_attention_geometry.py rename to tests/regional_generation/unet/test_unet_attention_geometry.py diff --git a/tests/test_unet_attention_phase_admission.py b/tests/regional_generation/unet/test_unet_attention_phase_admission.py similarity index 100% rename from tests/test_unet_attention_phase_admission.py rename to tests/regional_generation/unet/test_unet_attention_phase_admission.py diff --git a/tests/test_unet_attention_phase_diagnostics.py b/tests/regional_generation/unet/test_unet_attention_phase_diagnostics.py similarity index 100% rename from tests/test_unet_attention_phase_diagnostics.py rename to tests/regional_generation/unet/test_unet_attention_phase_diagnostics.py diff --git a/tests/test_unet_attention_phase_schedule.py b/tests/regional_generation/unet/test_unet_attention_phase_schedule.py similarity index 100% rename from tests/test_unet_attention_phase_schedule.py rename to tests/regional_generation/unet/test_unet_attention_phase_schedule.py diff --git a/tests/test_unet_attention_phase_session.py b/tests/regional_generation/unet/test_unet_attention_phase_session.py similarity index 100% rename from tests/test_unet_attention_phase_session.py rename to tests/regional_generation/unet/test_unet_attention_phase_session.py diff --git a/tests/test_unet_attn2_comfy_contract.py b/tests/regional_generation/unet/test_unet_attn2_comfy_contract.py similarity index 100% rename from tests/test_unet_attn2_comfy_contract.py rename to tests/regional_generation/unet/test_unet_attn2_comfy_contract.py diff --git a/tests/test_unet_attn2_execution.py b/tests/regional_generation/unet/test_unet_attn2_execution.py similarity index 100% rename from tests/test_unet_attn2_execution.py rename to tests/regional_generation/unet/test_unet_attn2_execution.py diff --git a/tests/test_unet_attn2_patch.py b/tests/regional_generation/unet/test_unet_attn2_patch.py similarity index 100% rename from tests/test_unet_attn2_patch.py rename to tests/regional_generation/unet/test_unet_attn2_patch.py diff --git a/tests/test_unet_attn2_resolution_cache.py b/tests/regional_generation/unet/test_unet_attn2_resolution_cache.py similarity index 100% rename from tests/test_unet_attn2_resolution_cache.py rename to tests/regional_generation/unet/test_unet_attn2_resolution_cache.py diff --git a/tests/test_unet_attn2_resolution_contract.py b/tests/regional_generation/unet/test_unet_attn2_resolution_contract.py similarity index 100% rename from tests/test_unet_attn2_resolution_contract.py rename to tests/regional_generation/unet/test_unet_attn2_resolution_contract.py diff --git a/tests/test_unet_dynamic_attention_integration.py b/tests/regional_generation/unet/test_unet_dynamic_attention_integration.py similarity index 100% rename from tests/test_unet_dynamic_attention_integration.py rename to tests/regional_generation/unet/test_unet_dynamic_attention_integration.py diff --git a/tests/test_unet_dynamic_attn2_execution_resolver.py b/tests/regional_generation/unet/test_unet_dynamic_attn2_execution_resolver.py similarity index 100% rename from tests/test_unet_dynamic_attn2_execution_resolver.py rename to tests/regional_generation/unet/test_unet_dynamic_attn2_execution_resolver.py diff --git a/tests/test_unet_regional_row_activity.py b/tests/regional_generation/unet/test_unet_regional_row_activity.py similarity index 100% rename from tests/test_unet_regional_row_activity.py rename to tests/regional_generation/unet/test_unet_regional_row_activity.py diff --git a/tests/sampling/__init__.py b/tests/sampling/__init__.py new file mode 100644 index 0000000..49758bc --- /dev/null +++ b/tests/sampling/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own sampling test behavior.""" diff --git a/tests/test_a1111_sampling.py b/tests/sampling/test_a1111_sampling.py similarity index 100% rename from tests/test_a1111_sampling.py rename to tests/sampling/test_a1111_sampling.py diff --git a/tests/test_ksampler_attention_coupling_v3_node.py b/tests/sampling/test_ksampler_attention_coupling_v3_node.py similarity index 100% rename from tests/test_ksampler_attention_coupling_v3_node.py rename to tests/sampling/test_ksampler_attention_coupling_v3_node.py diff --git a/tests/test_ksampler_contextual_attention_coupling_v3_node.py b/tests/sampling/test_ksampler_contextual_attention_coupling_v3_node.py similarity index 100% rename from tests/test_ksampler_contextual_attention_coupling_v3_node.py rename to tests/sampling/test_ksampler_contextual_attention_coupling_v3_node.py diff --git a/tests/test_ksampler_contextual_diffusion_node.py b/tests/sampling/test_ksampler_contextual_diffusion_node.py similarity index 100% rename from tests/test_ksampler_contextual_diffusion_node.py rename to tests/sampling/test_ksampler_contextual_diffusion_node.py diff --git a/tests/test_ksampler_extras_node.py b/tests/sampling/test_ksampler_extras_node.py similarity index 100% rename from tests/test_ksampler_extras_node.py rename to tests/sampling/test_ksampler_extras_node.py diff --git a/tests/test_ksampler_tiled_attention_coupling_v3_node.py b/tests/sampling/test_ksampler_tiled_attention_coupling_v3_node.py similarity index 100% rename from tests/test_ksampler_tiled_attention_coupling_v3_node.py rename to tests/sampling/test_ksampler_tiled_attention_coupling_v3_node.py diff --git a/tests/test_ksampler_tiled_diffusion_node.py b/tests/sampling/test_ksampler_tiled_diffusion_node.py similarity index 100% rename from tests/test_ksampler_tiled_diffusion_node.py rename to tests/sampling/test_ksampler_tiled_diffusion_node.py diff --git a/tests/test_latent_geometry.py b/tests/sampling/test_latent_geometry.py similarity index 100% rename from tests/test_latent_geometry.py rename to tests/sampling/test_latent_geometry.py diff --git a/tests/test_mixture_of_diffusers_sampling.py b/tests/sampling/test_mixture_of_diffusers_sampling.py similarity index 57% rename from tests/test_mixture_of_diffusers_sampling.py rename to tests/sampling/test_mixture_of_diffusers_sampling.py index d452eeb..0be0bba 100644 --- a/tests/test_mixture_of_diffusers_sampling.py +++ b/tests/sampling/test_mixture_of_diffusers_sampling.py @@ -14,7 +14,6 @@ import torch from simple_syrup.domain.segs import CropRegion from simple_syrup.domain.tiled_diffusion import gaussian_tile_weights from simple_syrup.runtime import mixture_of_diffusers_sampling as mod_sampling -from simple_syrup.runtime import sampling_samplers, sampling_schedulers from simple_syrup.runtime.detail_previews import DetailPreviewContext comfy_sample = mod_sampling._comfy_sample() @@ -392,284 +391,3 @@ def test_model_wrapper_delegates_shape_mismatch_unchanged() -> None: assert calls == 1 assert torch.allclose(output, x + 5.0) - - -def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Sampling mirrors KSampler flow while using a wrapped model clone.""" - - calls: dict[str, Any] = {} - model = FakeModel() - sampler = FakeSampler() - latent_samples = torch.zeros((1, 4, 4, 8), dtype=torch.float32) - fixed_noise = torch.ones_like(latent_samples) - fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) - sampled = torch.full_like(latent_samples, 0.25) - latent_image: dict[str, Any] = { - "samples": latent_samples, - "downscale_ratio_spacial": 2, - "kept": "value", - } - - def fake_resolve_sampler(sampler_name: str) -> FakeSampler: - """Record sampler resolution.""" - - calls["sampler_name"] = sampler_name - return sampler - - def fake_calculate_sigmas(**kwargs: object) -> torch.Tensor: - """Record scheduler calculation.""" - - calls["calculate_sigmas"] = kwargs - return fixed_sigmas - - def fake_fix_empty_latent_channels( - received_model: FakeModel, - samples: torch.Tensor, - downscale_ratio_spacial: int | None, - ) -> torch.Tensor: - """Record latent channel normalization.""" - - calls["fix_empty_latent_channels"] = { - "model": received_model, - "samples": samples, - "downscale_ratio_spacial": downscale_ratio_spacial, - } - return samples - - def fake_prepare_noise( - samples: torch.Tensor, - seed: int, - batch_inds: object = None, - ) -> torch.Tensor: - """Record noise preparation.""" - - calls["prepare_noise"] = { - "samples": samples, - "seed": seed, - "batch_inds": batch_inds, - } - return fixed_noise - - def fake_prepare_callback(received_model: FakeModel, steps: int) -> str: - """Record callback preparation.""" - - calls["prepare_callback"] = {"model": received_model, "steps": steps} - return "callback" - - monkeypatch.setattr(sampling_samplers, "resolve_sampler", fake_resolve_sampler) - monkeypatch.setattr( - sampling_schedulers, - "calculate_sigmas", - fake_calculate_sigmas, - ) - monkeypatch.setattr( - comfy_sample, - "fix_empty_latent_channels", - fake_fix_empty_latent_channels, - ) - monkeypatch.setattr(comfy_sample, "prepare_noise", fake_prepare_noise) - monkeypatch.setattr(latent_preview, "prepare_callback", fake_prepare_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: object, - disable_pbar: bool, - seed: int, - ) -> torch.Tensor: - """Record sample_custom arguments.""" - - calls["sample_custom"] = { - "model": received_model, - "noise": noise, - "cfg": cfg, - "sampler": received_sampler, - "sigmas": sigmas, - "positive": positive, - "negative": negative, - "latent_image": latent_image, - "noise_mask": noise_mask, - "callback": callback, - "disable_pbar": disable_pbar, - "seed": seed, - } - return sampled - - monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) - monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) - - output = mod_sampling.sample_mixture_of_diffusers( - model=model, - seed=123, - steps=2, - cfg=7.0, - sampler_name="euler", - scheduler="normal", - positive=[{"model_conds": {}}], - negative=[{"model_conds": {}}], - latent_image=latent_image, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=2, - ) - - assert output is not latent_image - assert output["samples"] is sampled - assert output["kept"] == "value" - assert "downscale_ratio_spacial" not in output - assert calls["sample_custom"]["model"] is not model - assert calls["sample_custom"]["model"].wrapper is not None - assert calls["sample_custom"]["sampler"] is sampler - assert calls["sample_custom"]["sigmas"] is fixed_sigmas - assert calls["sample_custom"]["disable_pbar"] is True - assert calls["calculate_sigmas"]["view"] == ( - sampling_schedulers.SchedulerView(latent_width=4, latent_height=4) - ) - - -def test_sample_accepts_singleton_depth_5d_latent( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Anima-style singleton-depth latents pass runtime validation.""" - - calls: dict[str, Any] = {} - model = FakeModel() - sampler = FakeSampler() - latent_samples = torch.zeros((1, 16, 1, 4, 8), dtype=torch.float32) - fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) - - 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=None: torch.ones_like(samples), - ) - monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) - monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", True) - - 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: object, - disable_pbar: bool, - seed: int, - ) -> torch.Tensor: - """Record sample_custom arguments and return a 5D latent.""" - - del noise, cfg, received_sampler, sigmas, positive, negative, noise_mask - del callback, disable_pbar, seed - calls["model"] = received_model - calls["latent_image"] = latent_image - return latent_image + 1.0 - - monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) - - output = mod_sampling.sample_mixture_of_diffusers( - model=model, - seed=123, - steps=2, - cfg=7.0, - sampler_name="euler", - scheduler="normal", - positive=[{"model_conds": {}}], - negative=[{"model_conds": {}}], - latent_image={"samples": latent_samples}, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=2, - ) - - assert calls["model"].wrapper is not None - assert calls["latent_image"] is latent_samples - assert torch.equal(output["samples"], latent_samples + 1.0) - - -def test_sample_rejects_5d_latent_with_non_singleton_depth( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Non-singleton 5D latents remain unsupported until validated explicitly.""" - - monkeypatch.setattr( - sampling_samplers, - "resolve_sampler", - lambda _sampler_name: FakeSampler(), - ) - monkeypatch.setattr( - sampling_schedulers, - "calculate_sigmas", - lambda **_kwargs: torch.tensor([1.0, 0.0], dtype=torch.float32), - ) - - with pytest.raises(ValueError, match="singleton third axis"): - mod_sampling.sample_mixture_of_diffusers( - model=FakeModel(), - seed=1, - steps=1, - cfg=1.0, - sampler_name="euler", - scheduler="normal", - positive=[], - negative=[], - latent_image={"samples": torch.zeros((1, 16, 2, 4, 4))}, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=1, - ) - - -def test_sample_rejects_unsupported_conditioning() -> None: - """Regional and ControlNet conditioning fail closed in the first slice.""" - - with pytest.raises(ValueError, match="regional conditioning or ControlNet"): - mod_sampling.sample_mixture_of_diffusers( - model=FakeModel(), - seed=1, - steps=1, - cfg=1.0, - sampler_name="euler", - scheduler="normal", - positive=[{"area": (4, 4, 0, 0)}], - negative=[], - latent_image={"samples": torch.zeros((1, 4, 4, 4))}, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=1, - ) diff --git a/tests/sampling/test_mixture_of_diffusers_sampling_runtime.py b/tests/sampling/test_mixture_of_diffusers_sampling_runtime.py new file mode 100644 index 0000000..c1b10e7 --- /dev/null +++ b/tests/sampling/test_mixture_of_diffusers_sampling_runtime.py @@ -0,0 +1,351 @@ +# 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 Mixture of Diffusers ComfyUI sampling runtime.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.runtime import mixture_of_diffusers_sampling as mod_sampling +from simple_syrup.runtime import sampling_samplers, sampling_schedulers + +comfy_sample = mod_sampling._comfy_sample() +comfy_utils = mod_sampling._comfy_utils() +latent_preview = mod_sampling._latent_preview() + + +class FakeModel: + """Provide the ModelPatcher methods used by the runtime.""" + + def __init__( + self, + model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, + ) -> None: + """Create a fake model patcher.""" + + self.load_device = torch.device("cpu") + 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.""" + + 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.""" + + 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.""" + + assert name == "model_sampling" + return self.model_sampling + + +class FakeSampler: + """Represent a resolved sampler in tests.""" + + def sample(self, *args: object, **kwargs: object) -> object: + """Provide ComfyUI's sampler protocol.""" + + del args, kwargs + return None + + +def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Sampling mirrors KSampler flow while using a wrapped model clone.""" + + calls: dict[str, Any] = {} + model = FakeModel() + sampler = FakeSampler() + latent_samples = torch.zeros((1, 4, 4, 8), dtype=torch.float32) + fixed_noise = torch.ones_like(latent_samples) + fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) + sampled = torch.full_like(latent_samples, 0.25) + latent_image: dict[str, Any] = { + "samples": latent_samples, + "downscale_ratio_spacial": 2, + "kept": "value", + } + + def fake_resolve_sampler(sampler_name: str) -> FakeSampler: + """Record sampler resolution.""" + + calls["sampler_name"] = sampler_name + return sampler + + def fake_calculate_sigmas(**kwargs: object) -> torch.Tensor: + """Record scheduler calculation.""" + + calls["calculate_sigmas"] = kwargs + return fixed_sigmas + + def fake_fix_empty_latent_channels( + received_model: FakeModel, + samples: torch.Tensor, + downscale_ratio_spacial: int | None, + ) -> torch.Tensor: + """Record latent channel normalization.""" + + calls["fix_empty_latent_channels"] = { + "model": received_model, + "samples": samples, + "downscale_ratio_spacial": downscale_ratio_spacial, + } + return samples + + def fake_prepare_noise( + samples: torch.Tensor, + seed: int, + batch_inds: object = None, + ) -> torch.Tensor: + """Record noise preparation.""" + + calls["prepare_noise"] = { + "samples": samples, + "seed": seed, + "batch_inds": batch_inds, + } + return fixed_noise + + def fake_prepare_callback(received_model: FakeModel, steps: int) -> str: + """Record callback preparation.""" + + calls["prepare_callback"] = {"model": received_model, "steps": steps} + return "callback" + + monkeypatch.setattr(sampling_samplers, "resolve_sampler", fake_resolve_sampler) + monkeypatch.setattr( + sampling_schedulers, + "calculate_sigmas", + fake_calculate_sigmas, + ) + monkeypatch.setattr( + comfy_sample, + "fix_empty_latent_channels", + fake_fix_empty_latent_channels, + ) + monkeypatch.setattr(comfy_sample, "prepare_noise", fake_prepare_noise) + monkeypatch.setattr(latent_preview, "prepare_callback", fake_prepare_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: object, + disable_pbar: bool, + seed: int, + ) -> torch.Tensor: + """Record sample_custom arguments.""" + + calls["sample_custom"] = { + "model": received_model, + "noise": noise, + "cfg": cfg, + "sampler": received_sampler, + "sigmas": sigmas, + "positive": positive, + "negative": negative, + "latent_image": latent_image, + "noise_mask": noise_mask, + "callback": callback, + "disable_pbar": disable_pbar, + "seed": seed, + } + return sampled + + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) + + output = mod_sampling.sample_mixture_of_diffusers( + model=model, + seed=123, + steps=2, + cfg=7.0, + sampler_name="euler", + scheduler="normal", + positive=[{"model_conds": {}}], + negative=[{"model_conds": {}}], + latent_image=latent_image, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=2, + ) + + assert output is not latent_image + assert output["samples"] is sampled + assert output["kept"] == "value" + assert "downscale_ratio_spacial" not in output + assert calls["sample_custom"]["model"] is not model + assert calls["sample_custom"]["model"].wrapper is not None + assert calls["sample_custom"]["sampler"] is sampler + assert calls["sample_custom"]["sigmas"] is fixed_sigmas + assert calls["sample_custom"]["disable_pbar"] is True + assert calls["calculate_sigmas"]["view"] == ( + sampling_schedulers.SchedulerView(latent_width=4, latent_height=4) + ) + + +def test_sample_accepts_singleton_depth_5d_latent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Anima-style singleton-depth latents pass runtime validation.""" + + calls: dict[str, Any] = {} + model = FakeModel() + sampler = FakeSampler() + latent_samples = torch.zeros((1, 16, 1, 4, 8), dtype=torch.float32) + fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) + + 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=None: torch.ones_like(samples), + ) + monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", True) + + 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: object, + disable_pbar: bool, + seed: int, + ) -> torch.Tensor: + """Record sample_custom arguments and return a 5D latent.""" + + del noise, cfg, received_sampler, sigmas, positive, negative, noise_mask + del callback, disable_pbar, seed + calls["model"] = received_model + calls["latent_image"] = latent_image + return latent_image + 1.0 + + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + + output = mod_sampling.sample_mixture_of_diffusers( + model=model, + seed=123, + steps=2, + cfg=7.0, + sampler_name="euler", + scheduler="normal", + positive=[{"model_conds": {}}], + negative=[{"model_conds": {}}], + latent_image={"samples": latent_samples}, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=2, + ) + + assert calls["model"].wrapper is not None + assert calls["latent_image"] is latent_samples + assert torch.equal(output["samples"], latent_samples + 1.0) + + +def test_sample_rejects_5d_latent_with_non_singleton_depth( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Non-singleton 5D latents remain unsupported until validated explicitly.""" + + monkeypatch.setattr( + sampling_samplers, + "resolve_sampler", + lambda _sampler_name: FakeSampler(), + ) + monkeypatch.setattr( + sampling_schedulers, + "calculate_sigmas", + lambda **_kwargs: torch.tensor([1.0, 0.0], dtype=torch.float32), + ) + + with pytest.raises(ValueError, match="singleton third axis"): + mod_sampling.sample_mixture_of_diffusers( + model=FakeModel(), + seed=1, + steps=1, + cfg=1.0, + sampler_name="euler", + scheduler="normal", + positive=[], + negative=[], + latent_image={"samples": torch.zeros((1, 16, 2, 4, 4))}, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=1, + ) + + +def test_sample_rejects_unsupported_conditioning() -> None: + """Regional and ControlNet conditioning fail closed in the first slice.""" + + with pytest.raises(ValueError, match="regional conditioning or ControlNet"): + mod_sampling.sample_mixture_of_diffusers( + model=FakeModel(), + seed=1, + steps=1, + cfg=1.0, + sampler_name="euler", + scheduler="normal", + positive=[{"area": (4, 4, 0, 0)}], + negative=[], + latent_image={"samples": torch.zeros((1, 4, 4, 4))}, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=1, + ) diff --git a/tests/test_multidiffusion_sampling.py b/tests/sampling/test_multidiffusion_sampling.py similarity index 68% rename from tests/test_multidiffusion_sampling.py rename to tests/sampling/test_multidiffusion_sampling.py index 194b219..fe58ea9 100644 --- a/tests/test_multidiffusion_sampling.py +++ b/tests/sampling/test_multidiffusion_sampling.py @@ -15,7 +15,6 @@ from simple_syrup.domain.segs import CropRegion from simple_syrup.runtime import ( multidiffusion_sampling, sampling_samplers, - sampling_schedulers, ) from simple_syrup.runtime.detail_previews import DetailPreviewContext @@ -568,253 +567,3 @@ def test_sample_rejects_unipc_before_sampler_resolution( latent_tile_overlap=0, latent_tile_batch_size=1, ) - - -def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Sampling mirrors KSampler flow while using a wrapped model clone.""" - - calls: dict[str, Any] = {} - model = FakeModel() - sampler = FakeSampler() - latent_samples = torch.zeros((1, 4, 4, 8), dtype=torch.float32) - fixed_noise = torch.ones_like(latent_samples) - fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) - sampled = torch.full_like(latent_samples, 0.25) - latent_image: dict[str, Any] = { - "samples": latent_samples, - "downscale_ratio_spacial": 2, - "kept": "value", - } - - monkeypatch.setattr( - sampling_samplers, - "resolve_sampler", - lambda _sampler_name: sampler, - ) - - def fake_calculate_sigmas(**kwargs: object) -> torch.Tensor: - """Record scheduler calculation.""" - - calls["calculate_sigmas"] = kwargs - return fixed_sigmas - - monkeypatch.setattr( - sampling_schedulers, - "calculate_sigmas", - fake_calculate_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=None: fixed_noise, - ) - monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) - monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) - - 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: object, - disable_pbar: bool, - seed: int, - ) -> torch.Tensor: - """Record sample_custom arguments.""" - - calls["sample_custom"] = { - "model": received_model, - "noise": noise, - "cfg": cfg, - "sampler": received_sampler, - "sigmas": sigmas, - "positive": positive, - "negative": negative, - "latent_image": latent_image, - "noise_mask": noise_mask, - "callback": callback, - "disable_pbar": disable_pbar, - "seed": seed, - } - return sampled - - monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) - - output = multidiffusion_sampling.sample_multidiffusion( - model=model, - seed=123, - steps=2, - cfg=7.0, - sampler_name="euler", - scheduler="normal", - positive=[{"model_conds": {}}], - negative=[{"model_conds": {}}], - latent_image=latent_image, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=2, - ) - - assert output is not latent_image - assert output["samples"] is sampled - assert output["kept"] == "value" - assert "downscale_ratio_spacial" not in output - assert calls["sample_custom"]["model"] is not model - assert calls["sample_custom"]["model"].wrapper is not None - assert calls["sample_custom"]["sampler"] is sampler - assert calls["sample_custom"]["noise"] is fixed_noise - assert calls["sample_custom"]["disable_pbar"] is True - assert calls["calculate_sigmas"]["view"] == ( - sampling_schedulers.SchedulerView(latent_width=4, latent_height=4) - ) - - -def test_sample_accepts_singleton_depth_5d_latent( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Anima-style singleton-depth latents pass runtime validation.""" - - calls: dict[str, Any] = {} - model = FakeModel() - sampler = FakeSampler() - latent_samples = torch.zeros((1, 16, 1, 4, 8), dtype=torch.float32) - fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) - - 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=None: torch.ones_like(samples), - ) - monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) - monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", True) - - 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: object, - disable_pbar: bool, - seed: int, - ) -> torch.Tensor: - """Record sample_custom arguments and return a 5D latent.""" - - del noise, cfg, received_sampler, sigmas, positive, negative, noise_mask - del callback, disable_pbar, seed - calls["model"] = received_model - calls["latent_image"] = latent_image - return latent_image + 1.0 - - monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) - - output = multidiffusion_sampling.sample_multidiffusion( - model=model, - seed=123, - steps=2, - cfg=7.0, - sampler_name="euler", - scheduler="normal", - positive=[{"model_conds": {}}], - negative=[{"model_conds": {}}], - latent_image={"samples": latent_samples}, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=2, - ) - - assert calls["model"].wrapper is not None - assert calls["latent_image"] is latent_samples - assert torch.equal(output["samples"], latent_samples + 1.0) - - -def test_sample_rejects_5d_latent_with_non_singleton_depth( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Non-singleton 5D latents remain unsupported until validated explicitly.""" - - monkeypatch.setattr( - sampling_samplers, - "resolve_sampler", - lambda _sampler_name: FakeSampler(), - ) - monkeypatch.setattr( - sampling_schedulers, - "calculate_sigmas", - lambda **_kwargs: torch.tensor([1.0, 0.0], dtype=torch.float32), - ) - - with pytest.raises(ValueError, match="singleton third axis"): - multidiffusion_sampling.sample_multidiffusion( - model=FakeModel(), - seed=1, - steps=1, - cfg=1.0, - sampler_name="euler", - scheduler="normal", - positive=[], - negative=[], - latent_image={"samples": torch.zeros((1, 16, 2, 4, 4))}, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=1, - ) - - -def test_sample_rejects_unsupported_conditioning() -> None: - """Regional and ControlNet conditioning fail closed in the first slice.""" - - with pytest.raises(ValueError, match="regional conditioning or ControlNet"): - multidiffusion_sampling.sample_multidiffusion( - model=FakeModel(), - seed=1, - steps=1, - cfg=1.0, - sampler_name="euler", - scheduler="normal", - positive=[{"area": (4, 4, 0, 0)}], - negative=[], - latent_image={"samples": torch.zeros((1, 4, 4, 4))}, - denoise=1.0, - latent_tile_width=4, - latent_tile_height=4, - latent_tile_overlap=0, - latent_tile_batch_size=1, - ) diff --git a/tests/sampling/test_multidiffusion_sampling_runtime.py b/tests/sampling/test_multidiffusion_sampling_runtime.py new file mode 100644 index 0000000..512275c --- /dev/null +++ b/tests/sampling/test_multidiffusion_sampling_runtime.py @@ -0,0 +1,323 @@ +# 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 MultiDiffusion ComfyUI sampling runtime.""" + +from __future__ import annotations + +from typing import Any + +import pytest +import torch + +from simple_syrup.runtime import ( + multidiffusion_sampling, + sampling_samplers, + sampling_schedulers, +) + +comfy_sample = multidiffusion_sampling._comfy_sample() +comfy_utils = multidiffusion_sampling._comfy_utils() +latent_preview = multidiffusion_sampling._latent_preview() + + +class FakeModel: + """Provide the ModelPatcher methods used by the runtime.""" + + def __init__( + self, + model_options: dict[str, Any] | None = None, + parent: FakeModel | None = None, + ) -> None: + """Create a fake model patcher.""" + + self.load_device = torch.device("cpu") + 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.""" + + 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.""" + + 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.""" + + assert name == "model_sampling" + return self.model_sampling + + +class FakeSampler: + """Represent a resolved sampler in tests.""" + + def sample(self, *args: object, **kwargs: object) -> object: + """Provide ComfyUI's sampler protocol.""" + + del args, kwargs + return None + + +def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Sampling mirrors KSampler flow while using a wrapped model clone.""" + + calls: dict[str, Any] = {} + model = FakeModel() + sampler = FakeSampler() + latent_samples = torch.zeros((1, 4, 4, 8), dtype=torch.float32) + fixed_noise = torch.ones_like(latent_samples) + fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) + sampled = torch.full_like(latent_samples, 0.25) + latent_image: dict[str, Any] = { + "samples": latent_samples, + "downscale_ratio_spacial": 2, + "kept": "value", + } + + monkeypatch.setattr( + sampling_samplers, + "resolve_sampler", + lambda _sampler_name: sampler, + ) + + def fake_calculate_sigmas(**kwargs: object) -> torch.Tensor: + """Record scheduler calculation.""" + + calls["calculate_sigmas"] = kwargs + return fixed_sigmas + + monkeypatch.setattr( + sampling_schedulers, + "calculate_sigmas", + fake_calculate_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=None: fixed_noise, + ) + monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) + + 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: object, + disable_pbar: bool, + seed: int, + ) -> torch.Tensor: + """Record sample_custom arguments.""" + + calls["sample_custom"] = { + "model": received_model, + "noise": noise, + "cfg": cfg, + "sampler": received_sampler, + "sigmas": sigmas, + "positive": positive, + "negative": negative, + "latent_image": latent_image, + "noise_mask": noise_mask, + "callback": callback, + "disable_pbar": disable_pbar, + "seed": seed, + } + return sampled + + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + + output = multidiffusion_sampling.sample_multidiffusion( + model=model, + seed=123, + steps=2, + cfg=7.0, + sampler_name="euler", + scheduler="normal", + positive=[{"model_conds": {}}], + negative=[{"model_conds": {}}], + latent_image=latent_image, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=2, + ) + + assert output is not latent_image + assert output["samples"] is sampled + assert output["kept"] == "value" + assert "downscale_ratio_spacial" not in output + assert calls["sample_custom"]["model"] is not model + assert calls["sample_custom"]["model"].wrapper is not None + assert calls["sample_custom"]["sampler"] is sampler + assert calls["sample_custom"]["noise"] is fixed_noise + assert calls["sample_custom"]["disable_pbar"] is True + assert calls["calculate_sigmas"]["view"] == ( + sampling_schedulers.SchedulerView(latent_width=4, latent_height=4) + ) + + +def test_sample_accepts_singleton_depth_5d_latent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Anima-style singleton-depth latents pass runtime validation.""" + + calls: dict[str, Any] = {} + model = FakeModel() + sampler = FakeSampler() + latent_samples = torch.zeros((1, 16, 1, 4, 8), dtype=torch.float32) + fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) + + 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=None: torch.ones_like(samples), + ) + monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) + monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", True) + + 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: object, + disable_pbar: bool, + seed: int, + ) -> torch.Tensor: + """Record sample_custom arguments and return a 5D latent.""" + + del noise, cfg, received_sampler, sigmas, positive, negative, noise_mask + del callback, disable_pbar, seed + calls["model"] = received_model + calls["latent_image"] = latent_image + return latent_image + 1.0 + + monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) + + output = multidiffusion_sampling.sample_multidiffusion( + model=model, + seed=123, + steps=2, + cfg=7.0, + sampler_name="euler", + scheduler="normal", + positive=[{"model_conds": {}}], + negative=[{"model_conds": {}}], + latent_image={"samples": latent_samples}, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=2, + ) + + assert calls["model"].wrapper is not None + assert calls["latent_image"] is latent_samples + assert torch.equal(output["samples"], latent_samples + 1.0) + + +def test_sample_rejects_5d_latent_with_non_singleton_depth( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Non-singleton 5D latents remain unsupported until validated explicitly.""" + + monkeypatch.setattr( + sampling_samplers, + "resolve_sampler", + lambda _sampler_name: FakeSampler(), + ) + monkeypatch.setattr( + sampling_schedulers, + "calculate_sigmas", + lambda **_kwargs: torch.tensor([1.0, 0.0], dtype=torch.float32), + ) + + with pytest.raises(ValueError, match="singleton third axis"): + multidiffusion_sampling.sample_multidiffusion( + model=FakeModel(), + seed=1, + steps=1, + cfg=1.0, + sampler_name="euler", + scheduler="normal", + positive=[], + negative=[], + latent_image={"samples": torch.zeros((1, 16, 2, 4, 4))}, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=1, + ) + + +def test_sample_rejects_unsupported_conditioning() -> None: + """Regional and ControlNet conditioning fail closed in the first slice.""" + + with pytest.raises(ValueError, match="regional conditioning or ControlNet"): + multidiffusion_sampling.sample_multidiffusion( + model=FakeModel(), + seed=1, + steps=1, + cfg=1.0, + sampler_name="euler", + scheduler="normal", + positive=[{"area": (4, 4, 0, 0)}], + negative=[], + latent_image={"samples": torch.zeros((1, 4, 4, 4))}, + denoise=1.0, + latent_tile_width=4, + latent_tile_height=4, + latent_tile_overlap=0, + latent_tile_batch_size=1, + ) diff --git a/tests/test_ordered_tensor_accumulation.py b/tests/sampling/test_ordered_tensor_accumulation.py similarity index 100% rename from tests/test_ordered_tensor_accumulation.py rename to tests/sampling/test_ordered_tensor_accumulation.py diff --git a/tests/test_phase_progress.py b/tests/sampling/test_phase_progress.py similarity index 100% rename from tests/test_phase_progress.py rename to tests/sampling/test_phase_progress.py diff --git a/tests/test_sampling_samplers.py b/tests/sampling/test_sampling_samplers.py similarity index 100% rename from tests/test_sampling_samplers.py rename to tests/sampling/test_sampling_samplers.py diff --git a/tests/sampling/test_sampling_scheduler_references.py b/tests/sampling/test_sampling_scheduler_references.py new file mode 100644 index 0000000..90a1621 --- /dev/null +++ b/tests/sampling/test_sampling_scheduler_references.py @@ -0,0 +1,449 @@ +# 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 KSampler Extras scheduler runtime helpers.""" + +from __future__ import annotations + +import math +from collections.abc import Sequence + +import comfy.samplers +import pytest +import torch + +from simple_syrup.runtime import sampling_schedulers +from simple_syrup.runtime.sampling_schedulers import ( + calculate_sigmas, +) + + +class FakeModel: + """Provide the model-sampling object expected by scheduler helpers.""" + + def __init__(self) -> None: + """Create a fake model with a stable model_sampling object.""" + + self.model_sampling = object() + + def get_model_object(self, name: str) -> object: + """Return the requested fake model object.""" + + assert name == "model_sampling" + return self.model_sampling + + +class FakeLatentFormat: + """Expose the spatial compression used to recover image dimensions.""" + + def __init__(self, spacial_downscale_ratio: int) -> None: + """Store one deterministic latent-to-image scale.""" + + self.spacial_downscale_ratio = spacial_downscale_ratio + + +class FakeModelWithLatentFormat(FakeModel): + """Provide model-sampling and latent-format objects for Flux2 tests.""" + + def __init__(self, spacial_downscale_ratio: int) -> None: + """Create a fake model with the requested spatial compression.""" + + super().__init__() + self.latent_format = FakeLatentFormat(spacial_downscale_ratio) + + def get_model_object(self, name: str) -> object: + """Return the requested fake model object.""" + + if name == "latent_format": + return self.latent_format + return super().get_model_object(name) + + +class FakeDiscreteModelSampling: + """Provide k-diffusion-style discrete sigma conversion for tests.""" + + def __init__(self, sigmas: Sequence[float]) -> None: + """Create a fake discrete model sampling object.""" + + self.sigmas = torch.tensor(sigmas, dtype=torch.float32) + self.log_sigmas = self.sigmas.log() + + def sigma(self, timestep: torch.Tensor) -> torch.Tensor: + """Convert fractional timesteps to sigmas with log-linear interpolation.""" + + timestep = torch.clamp( + timestep.float().to(self.log_sigmas.device), + min=0, + max=(len(self.sigmas) - 1), + ) + low_index = timestep.floor().long() + high_index = timestep.ceil().long() + weight = timestep.frac() + log_sigma = (1 - weight) * self.log_sigmas[ + low_index + ] + weight * self.log_sigmas[high_index] + return log_sigma.exp().to(timestep.device) + + +class FakeDiscreteModel: + """Provide a discrete model_sampling object for automatic_a1111 tests.""" + + def __init__(self) -> None: + """Create a fake model with k-diffusion-style sigmas.""" + + self.model_sampling = FakeDiscreteModelSampling((0.1, 0.3, 1.0)) + + def get_model_object(self, name: str) -> FakeDiscreteModelSampling: + """Return the requested fake model sampling object.""" + + assert name == "model_sampling" + return self.model_sampling + + +def assert_sigmas_close(actual: torch.Tensor, expected: list[float]) -> None: + """Assert that calculated sigmas match fixed reference values.""" + + assert torch.allclose( + actual.cpu(), + torch.tensor(expected, dtype=torch.float32), + atol=1e-5, + rtol=1e-5, + ) + + +def reference_extra_sigmas( + scheduler_name: str, + sampler_name: str, + steps: int, + denoise: float, +) -> torch.Tensor: + """Calculate expected extra scheduler sigmas with KSampler semantics.""" + + if denoise <= 0.0: + return torch.FloatTensor([]) + + schedule_steps = steps if denoise > 0.9999 else int(steps / denoise) + if sampler_name in comfy.samplers.KSampler.DISCARD_PENULTIMATE_SIGMA_SAMPLERS: + schedule_steps += 1 + + sigmas = reference_full_extra_schedule(scheduler_name, schedule_steps) + if sampler_name in comfy.samplers.KSampler.DISCARD_PENULTIMATE_SIGMA_SAMPLERS: + sigmas = torch.cat([sigmas[:-2], sigmas[-1:]]) + + if denoise <= 0.9999: + sigmas = sigmas[-(steps + 1) :] + return sigmas + + +def reference_full_extra_schedule(scheduler_name: str, steps: int) -> torch.Tensor: + """Calculate expected full AYS/GITS formula output for tests.""" + + if scheduler_name == "AYS SD1": + return reference_ays_schedule("SD1", steps) + if scheduler_name == "AYS SDXL": + return reference_ays_schedule("SDXL", steps) + if scheduler_name == "GITS": + return reference_gits_schedule(steps) + if scheduler_name == "automatic_a1111": + return reference_automatic_a1111_schedule(FakeDiscreteModel(), steps) + raise ValueError(f"Unsupported reference scheduler '{scheduler_name}'.") + + +def reference_ays_schedule(model_type: str, steps: int) -> torch.Tensor: + """Calculate full AYS schedule with Comfy Extras formula semantics.""" + + sigmas = list(sampling_schedulers.AYS_NOISE_LEVELS[model_type]) + if (steps + 1) != len(sigmas): + sigmas = reference_loglinear_interpolate(sigmas, steps + 1) + sigmas[-1] = 0.0 + return torch.FloatTensor(sigmas) + + +def reference_gits_schedule(steps: int) -> torch.Tensor: + """Calculate full GITS schedule for the default coefficient.""" + + if steps <= 20: + sigmas = list(sampling_schedulers.GITS_DEFAULT_NOISE_LEVELS[steps - 2]) + else: + sigmas = reference_loglinear_interpolate( + sampling_schedulers.GITS_DEFAULT_NOISE_LEVELS[-1], + steps + 1, + ) + sigmas[-1] = 0.0 + return torch.FloatTensor(sigmas) + + +def reference_automatic_a1111_schedule( + model: FakeDiscreteModel, + steps: int, +) -> torch.Tensor: + """Calculate k-diffusion DiscreteSchedule.get_sigmas-style output.""" + + model_sampling = model.model_sampling + timesteps = torch.linspace( + len(model_sampling.sigmas) - 1, + 0, + steps, + device=model_sampling.sigmas.device, + ) + sigmas = model_sampling.sigma(timesteps) + return torch.cat([sigmas, sigmas.new_zeros([1])]).cpu() + + +def reference_loglinear_interpolate( + sigmas: Sequence[float], + num_steps: int, +) -> list[float]: + """Interpolate reference sigma values in log space.""" + + reversed_logs = [math.log(value) for value in reversed(sigmas)] + source_max = len(reversed_logs) - 1 + target_max = num_steps - 1 + interpolated: list[float] = [] + + for target_index in range(num_steps): + source_position = target_index * source_max / target_max + left_index = math.floor(source_position) + right_index = min(left_index + 1, source_max) + fraction = source_position - left_index + left_value = reversed_logs[left_index] + right_value = reversed_logs[right_index] + interpolated.append( + math.exp(left_value + (right_value - left_value) * fraction) + ) + + return list(reversed(interpolated)) + + +def test_ays_sd1_full_schedule_matches_reference_values() -> None: + """AYS SD1 full schedules match fixed Comfy Extras reference values.""" + + model = FakeModel() + + assert_sigmas_close( + calculate_sigmas(model, "AYS SD1", "euler", 10, 1.0), + [ + 14.61464119, + 6.474576, + 3.86367464, + 2.69461513, + 1.88419211, + 1.39438045, + 0.96425837, + 0.65236861, + 0.39774564, + 0.15152326, + 0.0, + ], + ) + assert_sigmas_close( + calculate_sigmas(model, "AYS SD1", "euler", 20, 1.0), + [ + 14.61464119, + 9.72746658, + 6.474576, + 5.00156546, + 3.86367464, + 3.22662616, + 2.69461513, + 2.25325823, + 1.88419211, + 1.62088883, + 1.39438045, + 1.15954435, + 0.96425837, + 0.79312789, + 0.65236861, + 0.50938863, + 0.39774564, + 0.24549484, + 0.15152326, + 0.06647934, + 0.0, + ], + ) + + +def test_ays_sdxl_full_schedule_matches_reference_values() -> None: + """AYS SDXL full schedules match fixed Comfy Extras reference values.""" + + model = FakeModel() + + assert_sigmas_close( + calculate_sigmas(model, "AYS SDXL", "euler", 10, 1.0), + [ + 14.61464119, + 6.31844854, + 3.76817894, + 2.18114805, + 1.34052444, + 0.86207211, + 0.55506933, + 0.37985408, + 0.23323642, + 0.11141882, + 0.0, + ], + ) + assert_sigmas_close( + calculate_sigmas(model, "AYS SDXL", "euler", 20, 1.0), + [ + 14.61464119, + 9.60946751, + 6.31844854, + 4.87945127, + 3.76817894, + 2.86687231, + 2.18114805, + 1.70993638, + 1.34052444, + 1.07500172, + 0.86207211, + 0.69174403, + 0.55506933, + 0.45917898, + 0.37985408, + 0.29765046, + 0.23323642, + 0.16120461, + 0.11141882, + 0.05700676, + 0.0, + ], + ) + + +def test_gits_full_schedule_matches_reference_values() -> None: + """GITS full schedules match fixed reference values at the default coefficient.""" + + model = FakeModel() + + assert_sigmas_close( + calculate_sigmas(model, "GITS", "euler", 10, 1.0), + [ + 14.61464119, + 5.85520077, + 2.84484982, + 1.67050016, + 1.08895338, + 0.74807048, + 0.50118381, + 0.32104823, + 0.19894916, + 0.09824532, + 0.0, + ], + ) + assert_sigmas_close( + calculate_sigmas(model, "GITS", "euler", 20, 1.0), + [ + 14.61464119, + 7.49001646, + 4.65472794, + 3.07277966, + 2.19988537, + 1.61558151, + 1.24153244, + 0.95350921, + 0.74807048, + 0.59516323, + 0.50118381, + 0.41087446, + 0.34370604, + 0.29807833, + 0.25053367, + 0.22545385, + 0.19894916, + 0.17026083, + 0.13792117, + 0.09824532, + 0.0, + ], + ) + assert_sigmas_close( + calculate_sigmas(model, "GITS", "euler", 30, 1.0), + [ + 14.61464119, + 9.35946941, + 6.39175272, + 4.65472794, + 3.52900577, + 2.74887061, + 2.19988537, + 1.7906853, + 1.47980762, + 1.24153244, + 1.04120433, + 0.87942225, + 0.74807048, + 0.64230049, + 0.56202602, + 0.50118381, + 0.43900734, + 0.38714039, + 0.34370604, + 0.31257147, + 0.28130382, + 0.25053367, + 0.23352164, + 0.21624818, + 0.19894916, + 0.17933175, + 0.1587158, + 0.13792117, + 0.11000647, + 0.0655402, + 0.0, + ], + ) + + +@pytest.mark.parametrize("scheduler_name", ["AYS SD1", "AYS SDXL", "GITS"]) +@pytest.mark.parametrize("sampler_name", ["euler", "dpm_2"]) +@pytest.mark.parametrize("steps", [5, 10, 20, 30]) +@pytest.mark.parametrize("denoise", [0.25, 0.5, 0.8, 1.0]) +def test_extra_schedulers_follow_ksampler_denoise_semantics( + scheduler_name: str, + sampler_name: str, + steps: int, + denoise: float, +) -> None: + """Extra schedulers use KSampler partial-denoise and sigma cleanup rules.""" + + actual = calculate_sigmas( + FakeModel(), + scheduler_name, + sampler_name, + steps, + denoise, + ) + expected = reference_extra_sigmas(scheduler_name, sampler_name, steps, denoise) + + assert actual.shape == expected.shape + assert torch.allclose(actual, expected, atol=1e-5, rtol=1e-5) + if denoise < 1.0: + assert actual.shape == (steps + 1,) + + +def test_reported_ays_sd1_partial_denoise_regression() -> None: + """AYS SD1 with partial denoise returns a KSampler-length sigma schedule.""" + + actual = calculate_sigmas(FakeModel(), "AYS SD1", "euler", 5, 0.25) + expected = reference_extra_sigmas("AYS SD1", "euler", 5, 0.25) + + assert actual.shape == (6,) + assert torch.allclose(actual, expected, atol=1e-5, rtol=1e-5) + + +def test_gits_scope_uses_default_coefficient() -> None: + """The simple node contract exposes GITS at its default coefficient only.""" + + assert sampling_schedulers.GITS_DEFAULT_COEFF == 1.20 + + +def test_beta57_scope_uses_res4lyf_preset_parameters() -> None: + """The beta57 scheduler contract exposes RES4LYF's fixed beta preset.""" + + assert sampling_schedulers.BETA57_ALPHA == 0.5 + assert sampling_schedulers.BETA57_BETA == 0.7 diff --git a/tests/test_sampling_schedulers.py b/tests/sampling/test_sampling_schedulers.py similarity index 77% rename from tests/test_sampling_schedulers.py rename to tests/sampling/test_sampling_schedulers.py index 6eb4fd6..591a370 100644 --- a/tests/test_sampling_schedulers.py +++ b/tests/sampling/test_sampling_schedulers.py @@ -639,236 +639,3 @@ def test_automatic_a1111_rejects_unsupported_model_sampling_object() -> None: with pytest.raises(ValueError, match="automatic_a1111 requires"): calculate_sigmas(FakeModel(), "automatic_a1111", "euler", 4, 1.0) - - -def test_ays_sd1_full_schedule_matches_reference_values() -> None: - """AYS SD1 full schedules match fixed Comfy Extras reference values.""" - - model = FakeModel() - - assert_sigmas_close( - calculate_sigmas(model, "AYS SD1", "euler", 10, 1.0), - [ - 14.61464119, - 6.474576, - 3.86367464, - 2.69461513, - 1.88419211, - 1.39438045, - 0.96425837, - 0.65236861, - 0.39774564, - 0.15152326, - 0.0, - ], - ) - assert_sigmas_close( - calculate_sigmas(model, "AYS SD1", "euler", 20, 1.0), - [ - 14.61464119, - 9.72746658, - 6.474576, - 5.00156546, - 3.86367464, - 3.22662616, - 2.69461513, - 2.25325823, - 1.88419211, - 1.62088883, - 1.39438045, - 1.15954435, - 0.96425837, - 0.79312789, - 0.65236861, - 0.50938863, - 0.39774564, - 0.24549484, - 0.15152326, - 0.06647934, - 0.0, - ], - ) - - -def test_ays_sdxl_full_schedule_matches_reference_values() -> None: - """AYS SDXL full schedules match fixed Comfy Extras reference values.""" - - model = FakeModel() - - assert_sigmas_close( - calculate_sigmas(model, "AYS SDXL", "euler", 10, 1.0), - [ - 14.61464119, - 6.31844854, - 3.76817894, - 2.18114805, - 1.34052444, - 0.86207211, - 0.55506933, - 0.37985408, - 0.23323642, - 0.11141882, - 0.0, - ], - ) - assert_sigmas_close( - calculate_sigmas(model, "AYS SDXL", "euler", 20, 1.0), - [ - 14.61464119, - 9.60946751, - 6.31844854, - 4.87945127, - 3.76817894, - 2.86687231, - 2.18114805, - 1.70993638, - 1.34052444, - 1.07500172, - 0.86207211, - 0.69174403, - 0.55506933, - 0.45917898, - 0.37985408, - 0.29765046, - 0.23323642, - 0.16120461, - 0.11141882, - 0.05700676, - 0.0, - ], - ) - - -def test_gits_full_schedule_matches_reference_values() -> None: - """GITS full schedules match fixed reference values at the default coefficient.""" - - model = FakeModel() - - assert_sigmas_close( - calculate_sigmas(model, "GITS", "euler", 10, 1.0), - [ - 14.61464119, - 5.85520077, - 2.84484982, - 1.67050016, - 1.08895338, - 0.74807048, - 0.50118381, - 0.32104823, - 0.19894916, - 0.09824532, - 0.0, - ], - ) - assert_sigmas_close( - calculate_sigmas(model, "GITS", "euler", 20, 1.0), - [ - 14.61464119, - 7.49001646, - 4.65472794, - 3.07277966, - 2.19988537, - 1.61558151, - 1.24153244, - 0.95350921, - 0.74807048, - 0.59516323, - 0.50118381, - 0.41087446, - 0.34370604, - 0.29807833, - 0.25053367, - 0.22545385, - 0.19894916, - 0.17026083, - 0.13792117, - 0.09824532, - 0.0, - ], - ) - assert_sigmas_close( - calculate_sigmas(model, "GITS", "euler", 30, 1.0), - [ - 14.61464119, - 9.35946941, - 6.39175272, - 4.65472794, - 3.52900577, - 2.74887061, - 2.19988537, - 1.7906853, - 1.47980762, - 1.24153244, - 1.04120433, - 0.87942225, - 0.74807048, - 0.64230049, - 0.56202602, - 0.50118381, - 0.43900734, - 0.38714039, - 0.34370604, - 0.31257147, - 0.28130382, - 0.25053367, - 0.23352164, - 0.21624818, - 0.19894916, - 0.17933175, - 0.1587158, - 0.13792117, - 0.11000647, - 0.0655402, - 0.0, - ], - ) - - -@pytest.mark.parametrize("scheduler_name", ["AYS SD1", "AYS SDXL", "GITS"]) -@pytest.mark.parametrize("sampler_name", ["euler", "dpm_2"]) -@pytest.mark.parametrize("steps", [5, 10, 20, 30]) -@pytest.mark.parametrize("denoise", [0.25, 0.5, 0.8, 1.0]) -def test_extra_schedulers_follow_ksampler_denoise_semantics( - scheduler_name: str, - sampler_name: str, - steps: int, - denoise: float, -) -> None: - """Extra schedulers use KSampler partial-denoise and sigma cleanup rules.""" - - actual = calculate_sigmas( - FakeModel(), - scheduler_name, - sampler_name, - steps, - denoise, - ) - expected = reference_extra_sigmas(scheduler_name, sampler_name, steps, denoise) - - assert actual.shape == expected.shape - assert torch.allclose(actual, expected, atol=1e-5, rtol=1e-5) - if denoise < 1.0: - assert actual.shape == (steps + 1,) - - -def test_reported_ays_sd1_partial_denoise_regression() -> None: - """AYS SD1 with partial denoise returns a KSampler-length sigma schedule.""" - - actual = calculate_sigmas(FakeModel(), "AYS SD1", "euler", 5, 0.25) - expected = reference_extra_sigmas("AYS SD1", "euler", 5, 0.25) - - assert actual.shape == (6,) - assert torch.allclose(actual, expected, atol=1e-5, rtol=1e-5) - - -def test_gits_scope_uses_default_coefficient() -> None: - """The simple node contract exposes GITS at its default coefficient only.""" - - assert sampling_schedulers.GITS_DEFAULT_COEFF == 1.20 - - -def test_beta57_scope_uses_res4lyf_preset_parameters() -> None: - """The beta57 scheduler contract exposes RES4LYF's fixed beta preset.""" - - assert sampling_schedulers.BETA57_ALPHA == 0.5 - assert sampling_schedulers.BETA57_BETA == 0.7 diff --git a/tests/test_segs_tiled_diffusion.py b/tests/sampling/test_segs_tiled_diffusion.py similarity index 100% rename from tests/test_segs_tiled_diffusion.py rename to tests/sampling/test_segs_tiled_diffusion.py diff --git a/tests/test_tile_prediction_accumulation_characterization.py b/tests/sampling/test_tile_prediction_accumulation_characterization.py similarity index 100% rename from tests/test_tile_prediction_accumulation_characterization.py rename to tests/sampling/test_tile_prediction_accumulation_characterization.py diff --git a/tests/test_tiled_diffusion_conditioning_batch_service.py b/tests/sampling/test_tiled_diffusion_conditioning_batch_service.py similarity index 100% rename from tests/test_tiled_diffusion_conditioning_batch_service.py rename to tests/sampling/test_tiled_diffusion_conditioning_batch_service.py diff --git a/tests/test_tiled_diffusion_domain.py b/tests/sampling/test_tiled_diffusion_domain.py similarity index 100% rename from tests/test_tiled_diffusion_domain.py rename to tests/sampling/test_tiled_diffusion_domain.py diff --git a/tests/test_tiled_diffusion_item_sampling_service.py b/tests/sampling/test_tiled_diffusion_item_sampling_service.py similarity index 100% rename from tests/test_tiled_diffusion_item_sampling_service.py rename to tests/sampling/test_tiled_diffusion_item_sampling_service.py diff --git a/tests/test_tiled_diffusion_sampling_routing_service.py b/tests/sampling/test_tiled_diffusion_sampling_routing_service.py similarity index 100% rename from tests/test_tiled_diffusion_sampling_routing_service.py rename to tests/sampling/test_tiled_diffusion_sampling_routing_service.py diff --git a/tests/test_tiled_sampling_validation_characterization.py b/tests/sampling/test_tiled_sampling_validation_characterization.py similarity index 100% rename from tests/test_tiled_sampling_validation_characterization.py rename to tests/sampling/test_tiled_sampling_validation_characterization.py diff --git a/tests/test_tiled_spatial_batch_layout.py b/tests/sampling/test_tiled_spatial_batch_layout.py similarity index 100% rename from tests/test_tiled_spatial_batch_layout.py rename to tests/sampling/test_tiled_spatial_batch_layout.py diff --git a/tests/segmentation/__init__.py b/tests/segmentation/__init__.py new file mode 100644 index 0000000..ae16441 --- /dev/null +++ b/tests/segmentation/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own segmentation test behavior.""" diff --git a/tests/segmentation/detection/__init__.py b/tests/segmentation/detection/__init__.py new file mode 100644 index 0000000..a5304e6 --- /dev/null +++ b/tests/segmentation/detection/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own segmentation detection test behavior.""" diff --git a/tests/test_detect_segs_with_ultralytics_node.py b/tests/segmentation/detection/test_detect_segs_with_ultralytics_node.py similarity index 99% rename from tests/test_detect_segs_with_ultralytics_node.py rename to tests/segmentation/detection/test_detect_segs_with_ultralytics_node.py index f7b2ea1..6593b13 100644 --- a/tests/test_detect_segs_with_ultralytics_node.py +++ b/tests/segmentation/detection/test_detect_segs_with_ultralytics_node.py @@ -20,7 +20,7 @@ from simple_syrup.domain.segs import ( Segment, ) from simple_syrup.nodes.detect_segs_with_ultralytics import DetectSEGSWithUltralytics -from simple_syrup.runtime.ultralytics_loader import UltralyticsDetectorModel +from simple_syrup.runtime.ultralytics_model_adapter import UltralyticsDetectorModel from simple_syrup.services.segs_output_service import CombinedSegsResult diff --git a/tests/test_detector_compat.py b/tests/segmentation/detection/test_detector_compat.py similarity index 66% rename from tests/test_detector_compat.py rename to tests/segmentation/detection/test_detector_compat.py index 8851002..b276ba8 100644 --- a/tests/test_detector_compat.py +++ b/tests/segmentation/detection/test_detector_compat.py @@ -7,23 +7,20 @@ from __future__ import annotations from pathlib import Path -from typing import cast +from typing import Any, cast -import pytest import torch from simple_syrup.domain.segs import BoundingBox, CropRegion, Segment -from simple_syrup.runtime.detector_compat import BBoxDetectorFacade, SegmDetectorFacade -from simple_syrup.runtime.ultralytics_loader import UltralyticsDetectorModel +from simple_syrup.runtime.ultralytics_model_adapter import UltralyticsDetectorModel +from simple_syrup.services.detector_compat import BBoxDetectorFacade, SegmDetectorFacade -def test_bbox_facade_accepts_detector_signature( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_bbox_facade_accepts_detector_signature() -> None: """BBox facade accepts the expected detector arguments.""" - _patch_service(monkeypatch, prefer_segmentation_expected=False) - facade = BBoxDetectorFacade(_model(supports_segmentation=False)) + service = _FakeService(expected_prefer_segmentation=False) + facade = BBoxDetectorFacade(_model(supports_segmentation=False), cast(Any, service)) header, segments = cast( tuple[object, list[Segment]], @@ -34,28 +31,24 @@ def test_bbox_facade_accepts_detector_signature( assert segments[0].label == "face" -def test_segmentation_facade_has_bbox_detector( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_segmentation_facade_has_bbox_detector() -> None: """Segmentation facade exposes its paired bbox detector.""" - _patch_service(monkeypatch, prefer_segmentation_expected=True) - bbox = BBoxDetectorFacade(_model(supports_segmentation=True)) - facade = SegmDetectorFacade(bbox.detector_model, bbox) + service = _FakeService(expected_prefer_segmentation=True) + bbox = BBoxDetectorFacade(_model(supports_segmentation=True), cast(Any, service)) + facade = SegmDetectorFacade(bbox.detector_model, bbox, cast(Any, service)) facade.detect(torch.zeros((1, 8, 8, 3)), 0.5, 1, 2.0) assert facade.bbox_detector is bbox -def test_segmentation_facade_falls_back_for_bbox_model( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_segmentation_facade_falls_back_for_bbox_model() -> None: """A bbox-only native model remains safe on the segmentation facade.""" - _patch_service(monkeypatch, prefer_segmentation_expected=False) - bbox = BBoxDetectorFacade(_model(supports_segmentation=False)) - facade = SegmDetectorFacade(bbox.detector_model, bbox) + service = _FakeService(expected_prefer_segmentation=False) + bbox = BBoxDetectorFacade(_model(supports_segmentation=False), cast(Any, service)) + facade = SegmDetectorFacade(bbox.detector_model, bbox, cast(Any, service)) header, segments = cast( tuple[object, list[Segment]], @@ -69,7 +62,10 @@ def test_segmentation_facade_falls_back_for_bbox_model( class _FakeService: """Fake native detection service for facade tests.""" - expected_prefer_segmentation: bool = False + def __init__(self, expected_prefer_segmentation: bool) -> None: + """Store the expected detection mode.""" + + self.expected_prefer_segmentation = expected_prefer_segmentation def detect( self, @@ -96,18 +92,6 @@ class _FakeService: return (8, 8), (segment,) -def _patch_service( - monkeypatch: pytest.MonkeyPatch, - prefer_segmentation_expected: bool, -) -> None: - """Patch the lazily imported detection service used by facades.""" - - import simple_syrup.services.segs_detection_service as service_module - - _FakeService.expected_prefer_segmentation = prefer_segmentation_expected - monkeypatch.setattr(service_module, "SegsDetectionService", _FakeService) - - def _model(supports_segmentation: bool) -> UltralyticsDetectorModel: """Return a native detector model test double.""" diff --git a/tests/test_grounded_sam_model_info_node.py b/tests/segmentation/detection/test_grounded_sam_model_info_node.py similarity index 100% rename from tests/test_grounded_sam_model_info_node.py rename to tests/segmentation/detection/test_grounded_sam_model_info_node.py diff --git a/tests/test_grounding_dino_bert_adapter.py b/tests/segmentation/detection/test_grounding_dino_bert_adapter.py similarity index 100% rename from tests/test_grounding_dino_bert_adapter.py rename to tests/segmentation/detection/test_grounding_dino_bert_adapter.py diff --git a/tests/test_grounding_dino_loader.py b/tests/segmentation/detection/test_grounding_dino_loader.py similarity index 99% rename from tests/test_grounding_dino_loader.py rename to tests/segmentation/detection/test_grounding_dino_loader.py index 98cd07a..eeeba6b 100644 --- a/tests/test_grounding_dino_loader.py +++ b/tests/segmentation/detection/test_grounding_dino_loader.py @@ -15,6 +15,7 @@ from typing import Any import pytest import torch +from support.helpers import FakeFolderPaths from simple_syrup.runtime.grounding_dino_loader import ( GROUNDING_DINO_RUNTIME_PACKAGE, @@ -25,7 +26,6 @@ from simple_syrup.runtime.grounding_dino_loader import ( GroundingDINOModelCacheKey, ) from simple_syrup.runtime.loaded_models import LoadedGroundingDINOModel -from test_helpers import FakeFolderPaths @dataclass diff --git a/tests/test_grounding_dino_model_loader_node.py b/tests/segmentation/detection/test_grounding_dino_model_loader_node.py similarity index 100% rename from tests/test_grounding_dino_model_loader_node.py rename to tests/segmentation/detection/test_grounding_dino_model_loader_node.py diff --git a/tests/test_grounding_dino_text_token_masks.py b/tests/segmentation/detection/test_grounding_dino_text_token_masks.py similarity index 96% rename from tests/test_grounding_dino_text_token_masks.py rename to tests/segmentation/detection/test_grounding_dino_text_token_masks.py index 6c6f863..92f6534 100644 --- a/tests/test_grounding_dino_text_token_masks.py +++ b/tests/segmentation/detection/test_grounding_dino_text_token_masks.py @@ -13,8 +13,9 @@ from types import ModuleType from typing import cast import torch +from support.repository import REPOSITORY_ROOT -REPO_ROOT = Path(__file__).resolve().parents[1] +REPO_ROOT = REPOSITORY_ROOT TEXT_TOKEN_MASKS_PATH = ( REPO_ROOT / "simple_syrup" diff --git a/tests/test_layerstyle_sam_models_adapter_node.py b/tests/segmentation/detection/test_layerstyle_sam_models_adapter_node.py similarity index 100% rename from tests/test_layerstyle_sam_models_adapter_node.py rename to tests/segmentation/detection/test_layerstyle_sam_models_adapter_node.py diff --git a/tests/test_load_ultralytics_model_node.py b/tests/segmentation/detection/test_load_ultralytics_model_node.py similarity index 95% rename from tests/test_load_ultralytics_model_node.py rename to tests/segmentation/detection/test_load_ultralytics_model_node.py index 56a3a3c..9010e37 100644 --- a/tests/test_load_ultralytics_model_node.py +++ b/tests/segmentation/detection/test_load_ultralytics_model_node.py @@ -12,7 +12,7 @@ import pytest from simple_syrup.nodes.load_ultralytics_model import LoadUltralyticsModel from simple_syrup.runtime.model_downloads import ProgressReporter -from simple_syrup.runtime.ultralytics_loader import LoadedUltralyticsDetector +from simple_syrup.services.ultralytics_loader_service import LoadedUltralyticsDetector def test_load_ultralytics_model_node_contract( diff --git a/tests/test_sam_automatic_segmenter.py b/tests/segmentation/detection/test_sam_automatic_segmenter.py similarity index 100% rename from tests/test_sam_automatic_segmenter.py rename to tests/segmentation/detection/test_sam_automatic_segmenter.py diff --git a/tests/test_sam_loader.py b/tests/segmentation/detection/test_sam_loader.py similarity index 99% rename from tests/test_sam_loader.py rename to tests/segmentation/detection/test_sam_loader.py index 3343dd7..702a452 100644 --- a/tests/test_sam_loader.py +++ b/tests/segmentation/detection/test_sam_loader.py @@ -12,11 +12,11 @@ from pathlib import Path from types import ModuleType import pytest +from support.helpers import FakeFolderPaths from simple_syrup.runtime.loaded_models import LoadedSAMModel from simple_syrup.runtime.model_downloads import DownloadRequest, DownloadResult from simple_syrup.runtime.sam_loader import SAMLoaderService, SAMModelCacheKey -from test_helpers import FakeFolderPaths @dataclass diff --git a/tests/test_sam_model_loader_node.py b/tests/segmentation/detection/test_sam_model_loader_node.py similarity index 100% rename from tests/test_sam_model_loader_node.py rename to tests/segmentation/detection/test_sam_model_loader_node.py diff --git a/tests/test_sam_region_overlay_renderer.py b/tests/segmentation/detection/test_sam_region_overlay_renderer.py similarity index 100% rename from tests/test_sam_region_overlay_renderer.py rename to tests/segmentation/detection/test_sam_region_overlay_renderer.py diff --git a/tests/test_sam_segmenter.py b/tests/segmentation/detection/test_sam_segmenter.py similarity index 99% rename from tests/test_sam_segmenter.py rename to tests/segmentation/detection/test_sam_segmenter.py index 91b20b0..bc99d7d 100644 --- a/tests/test_sam_segmenter.py +++ b/tests/segmentation/detection/test_sam_segmenter.py @@ -11,11 +11,11 @@ from types import ModuleType, SimpleNamespace import pytest import torch +from support.helpers import make_image_tensor from simple_syrup.runtime.loaded_models import LoadedSAMModel from simple_syrup.runtime.model_device_manager import TorchModelDeviceManager from simple_syrup.runtime.sam_segmenter import SAMModelSegmenter -from test_helpers import make_image_tensor class RecordingWrapper: diff --git a/tests/test_tag_segs_with_external_llm_service.py b/tests/segmentation/detection/test_tag_segs_with_external_llm_service.py similarity index 100% rename from tests/test_tag_segs_with_external_llm_service.py rename to tests/segmentation/detection/test_tag_segs_with_external_llm_service.py diff --git a/tests/test_tag_segs_with_external_llm_v3_node.py b/tests/segmentation/detection/test_tag_segs_with_external_llm_v3_node.py similarity index 100% rename from tests/test_tag_segs_with_external_llm_v3_node.py rename to tests/segmentation/detection/test_tag_segs_with_external_llm_v3_node.py diff --git a/tests/test_tag_segs_with_wd14_node.py b/tests/segmentation/detection/test_tag_segs_with_wd14_node.py similarity index 100% rename from tests/test_tag_segs_with_wd14_node.py rename to tests/segmentation/detection/test_tag_segs_with_wd14_node.py diff --git a/tests/test_tag_segs_with_wd14_service.py b/tests/segmentation/detection/test_tag_segs_with_wd14_service.py similarity index 100% rename from tests/test_tag_segs_with_wd14_service.py rename to tests/segmentation/detection/test_tag_segs_with_wd14_service.py diff --git a/tests/test_tag_segs_with_wd14_v3_node.py b/tests/segmentation/detection/test_tag_segs_with_wd14_v3_node.py similarity index 100% rename from tests/test_tag_segs_with_wd14_v3_node.py rename to tests/segmentation/detection/test_tag_segs_with_wd14_v3_node.py diff --git a/tests/test_text_box_detector.py b/tests/segmentation/detection/test_text_box_detector.py similarity index 99% rename from tests/test_text_box_detector.py rename to tests/segmentation/detection/test_text_box_detector.py index 7fd7007..d128d5c 100644 --- a/tests/test_text_box_detector.py +++ b/tests/segmentation/detection/test_text_box_detector.py @@ -13,6 +13,7 @@ from typing import Any, cast import pytest import torch +from support.helpers import make_image_tensor from simple_syrup.domain.segs import BoundingBox from simple_syrup.runtime.grounding_dino_loader import GROUNDING_DINO_RUNTIME_PACKAGE @@ -22,7 +23,6 @@ from simple_syrup.runtime.text_box_detector import ( GroundingDINOTextBoxDetector, TextBoxDetection, ) -from test_helpers import make_image_tensor class PredictBoxesModel: diff --git a/tests/test_tile_and_tag_segs_node.py b/tests/segmentation/detection/test_tile_and_tag_segs_node.py similarity index 100% rename from tests/test_tile_and_tag_segs_node.py rename to tests/segmentation/detection/test_tile_and_tag_segs_node.py diff --git a/tests/test_tile_and_tag_segs_service.py b/tests/segmentation/detection/test_tile_and_tag_segs_service.py similarity index 100% rename from tests/test_tile_and_tag_segs_service.py rename to tests/segmentation/detection/test_tile_and_tag_segs_service.py diff --git a/tests/test_tile_and_tag_segs_v3_node.py b/tests/segmentation/detection/test_tile_and_tag_segs_v3_node.py similarity index 100% rename from tests/test_tile_and_tag_segs_v3_node.py rename to tests/segmentation/detection/test_tile_and_tag_segs_v3_node.py diff --git a/tests/test_ultralytics_detection_service.py b/tests/segmentation/detection/test_ultralytics_detection_service.py similarity index 74% rename from tests/test_ultralytics_detection_service.py rename to tests/segmentation/detection/test_ultralytics_detection_service.py index 7680b3f..b2bf01b 100644 --- a/tests/test_ultralytics_detection_service.py +++ b/tests/segmentation/detection/test_ultralytics_detection_service.py @@ -12,17 +12,16 @@ from typing import cast import pytest import torch -from simple_syrup.domain.segs import BoundingBox, CropRegion, NativeSegs, Segment +from simple_syrup.domain.segs import BoundingBox, CropRegion from simple_syrup.runtime.ultralytics_detection import ( UltralyticsDetection, parse_ultralytics_result, ) -from simple_syrup.runtime.ultralytics_loader import UltralyticsDetectorModel +from simple_syrup.runtime.ultralytics_model_adapter import UltralyticsDetectorModel from simple_syrup.services.segs_detection_service import ( DetectionRunner, SegsDetectionService, ) -from simple_syrup.services.segs_output_service import build_combined_segs_result def test_bbox_prediction_creates_rectangular_mask_segs() -> None: @@ -313,102 +312,6 @@ def test_empty_detections_return_empty_segs() -> None: assert service.detect(_image(), _model(False), 0.5, 0, 1.0, 1) == ((8, 8), ()) -def test_empty_segs_build_empty_combined_outputs() -> None: - """Empty SEGS produce empty combined SEGS and a zero mask.""" - - result = build_combined_segs_result(_image(), ((8, 8), ()), crop_factor=1.0) - - assert result.segs == ((8, 8), ()) - assert result.mask.shape == (1, 8, 8) - assert result.mask.dtype == torch.float32 - assert result.mask.sum().item() == 0.0 - - -def test_combined_result_unions_all_segment_masks() -> None: - """Combined SEGS contain one unioned segment for all source masks.""" - - image = torch.arange(1 * 8 * 8 * 3, dtype=torch.float32).reshape((1, 8, 8, 3)) - image = image / image.max() - first = _segment(CropRegion(1, 1, 3, 3), 0.4, torch.ones((2, 2))) - second = _segment( - CropRegion(4, 5, 7, 7), - 0.9, - [[2.0, 2.0, 2.0], [2.0, 2.0, 2.0]], - ) - - result = build_combined_segs_result(image, _segs(first, second), crop_factor=1.0) - _header, segments = result.segs - - assert result.mask.shape == (1, 8, 8) - assert result.mask.sum().item() == 10.0 - assert len(segments) == 1 - assert segments[0].crop_region == CropRegion(1, 1, 7, 7) - assert segments[0].bbox == BoundingBox(1, 1, 7, 7) - assert segments[0].label == "combined" - assert segments[0].confidence == 0.9 - assert cast(torch.Tensor, segments[0].cropped_mask).shape == (6, 6) - assert torch.equal( - cast(torch.Tensor, segments[0].cropped_image), image[:, 1:7, 1:7, :] - ) - - -def test_combined_result_applies_crop_factor_after_combining() -> None: - """Combined SEGS expand one unioned target for downstream detailers.""" - - image = torch.arange(1 * 8 * 8 * 3, dtype=torch.float32).reshape((1, 8, 8, 3)) - image = image / image.max() - first_mask = torch.zeros((2, 2), dtype=torch.float32) - first_mask[0, 0] = 1.0 - second_mask = torch.zeros((2, 2), dtype=torch.float32) - second_mask[1, 1] = 1.0 - first = Segment( - cropped_image=None, - cropped_mask=first_mask, - confidence=0.8, - crop_region=CropRegion(2, 2, 4, 4), - bbox=BoundingBox(2, 2, 3, 3), - label="face", - ) - second = Segment( - cropped_image=None, - cropped_mask=second_mask, - confidence=0.9, - crop_region=CropRegion(4, 4, 6, 6), - bbox=BoundingBox(5, 5, 6, 6), - label="face", - ) - - result = build_combined_segs_result(image, _segs(first, second), crop_factor=2.0) - _header, segments = result.segs - - assert len(segments) == 1 - assert segments[0].crop_region == CropRegion(0, 0, 8, 8) - assert segments[0].bbox == BoundingBox(2, 2, 6, 6) - assert cast(torch.Tensor, segments[0].cropped_mask).shape == (8, 8) - assert torch.equal(cast(torch.Tensor, segments[0].cropped_image), image) - - -def test_combined_mask_uses_max_for_overlaps() -> None: - """Overlapping source masks combine by max instead of addition.""" - - first = _segment(CropRegion(2, 2, 4, 4), 0.4, torch.full((2, 2), 0.7)) - second = _segment(CropRegion(2, 2, 4, 4), 0.8, torch.full((2, 2), 0.8)) - - result = build_combined_segs_result(_image(), _segs(first, second), crop_factor=1.0) - - assert result.mask[0, 2, 2].item() == pytest.approx(0.8) - assert result.mask.sum().item() == pytest.approx(3.2) - - -def test_combined_outputs_reject_mismatched_cropped_mask_shape() -> None: - """Invalid crop-local masks fail before producing misleading outputs.""" - - segment = _segment(CropRegion(1, 1, 3, 3), 1.0, torch.ones((3, 3))) - - with pytest.raises(ValueError, match="cropped_mask must match"): - build_combined_segs_result(_image(), _segs(segment), crop_factor=1.0) - - def test_batch_image_input_fails_clearly() -> None: """The first version rejects batched images.""" @@ -424,34 +327,6 @@ def _bbox_detection() -> UltralyticsDetection: return UltralyticsDetection(BoundingBox(2, 2, 5, 5), 0.9, "face", None) -def _segs(*segments: Segment) -> NativeSegs: - """Return native SEGS for an 8x8 image.""" - - return (8, 8), tuple(segments) - - -def _segment( - crop_region: CropRegion, - confidence: float, - cropped_mask: object, -) -> Segment: - """Return one segment for combined-output tests.""" - - return Segment( - cropped_image=None, - cropped_mask=cropped_mask, - confidence=confidence, - crop_region=crop_region, - bbox=BoundingBox( - crop_region.left, - crop_region.top, - crop_region.right, - crop_region.bottom, - ), - label="face", - ) - - def _runner( *detections: UltralyticsDetection, ) -> DetectionRunner: diff --git a/tests/test_ultralytics_loader.py b/tests/segmentation/detection/test_ultralytics_loader.py similarity index 99% rename from tests/test_ultralytics_loader.py rename to tests/segmentation/detection/test_ultralytics_loader.py index f2c6f38..43e0521 100644 --- a/tests/test_ultralytics_loader.py +++ b/tests/segmentation/detection/test_ultralytics_loader.py @@ -22,7 +22,7 @@ from simple_syrup.runtime.model_downloads import ( ProgressReporter, ) from simple_syrup.runtime.settings import SimpleSyrupSettings -from simple_syrup.runtime.ultralytics_loader import ( +from simple_syrup.services.ultralytics_loader_service import ( NO_LOCAL_ULTRALYTICS_MODELS, LoadedUltralyticsDetector, UltralyticsLoaderService, diff --git a/tests/test_vitmatte_loader.py b/tests/segmentation/detection/test_vitmatte_loader.py similarity index 99% rename from tests/test_vitmatte_loader.py rename to tests/segmentation/detection/test_vitmatte_loader.py index 6fe3998..0aecd71 100644 --- a/tests/test_vitmatte_loader.py +++ b/tests/segmentation/detection/test_vitmatte_loader.py @@ -12,6 +12,7 @@ from pathlib import Path from types import ModuleType import pytest +from support.helpers import FakeFolderPaths from simple_syrup.runtime.loaded_models import LoadedViTMatteModel from simple_syrup.runtime.model_catalog import get_vitmatte_entry, vitmatte_choices @@ -20,7 +21,6 @@ from simple_syrup.runtime.vitmatte_loader import ( ViTMatteModelCacheKey, is_valid_vitmatte_directory, ) -from test_helpers import FakeFolderPaths class RecordingSnapshotDownloader: diff --git a/tests/test_vitmatte_model_loader_node.py b/tests/segmentation/detection/test_vitmatte_model_loader_node.py similarity index 100% rename from tests/test_vitmatte_model_loader_node.py rename to tests/segmentation/detection/test_vitmatte_model_loader_node.py diff --git a/tests/test_vitmatte_refiner.py b/tests/segmentation/detection/test_vitmatte_refiner.py similarity index 98% rename from tests/test_vitmatte_refiner.py rename to tests/segmentation/detection/test_vitmatte_refiner.py index e66efc0..2c7afd0 100644 --- a/tests/test_vitmatte_refiner.py +++ b/tests/segmentation/detection/test_vitmatte_refiner.py @@ -12,6 +12,7 @@ from types import ModuleType, SimpleNamespace import pytest import torch +from support.helpers import make_image_tensor from simple_syrup.masking.mask_ops import MaskRefinementSettings from simple_syrup.runtime.loaded_models import LoadedViTMatteModel @@ -20,7 +21,6 @@ from simple_syrup.runtime.vitmatte_refiner import ( ViTMatteRefiner, generate_vitmatte_trimap, ) -from test_helpers import make_image_tensor class FakeProcessor: diff --git a/tests/test_wd14_tagger_loader.py b/tests/segmentation/detection/test_wd14_tagger_loader.py similarity index 99% rename from tests/test_wd14_tagger_loader.py rename to tests/segmentation/detection/test_wd14_tagger_loader.py index 488aa1c..4f7ca7a 100644 --- a/tests/test_wd14_tagger_loader.py +++ b/tests/segmentation/detection/test_wd14_tagger_loader.py @@ -11,6 +11,7 @@ from pathlib import Path from types import ModuleType import pytest +from support.helpers import FakeFolderPaths from simple_syrup.runtime.loaded_models import LoadedWD14Tagger from simple_syrup.runtime.model_downloads import ( @@ -23,7 +24,6 @@ from simple_syrup.runtime.wd14_tagger_loader import ( WD14TaggerCacheKey, WD14TaggerLoaderService, ) -from test_helpers import FakeFolderPaths def test_loader_reuses_existing_files_without_download( diff --git a/tests/test_wd14_tagger_loader_node.py b/tests/segmentation/detection/test_wd14_tagger_loader_node.py similarity index 100% rename from tests/test_wd14_tagger_loader_node.py rename to tests/segmentation/detection/test_wd14_tagger_loader_node.py diff --git a/tests/test_wd14_tagger_loader_v3_node.py b/tests/segmentation/detection/test_wd14_tagger_loader_v3_node.py similarity index 100% rename from tests/test_wd14_tagger_loader_v3_node.py rename to tests/segmentation/detection/test_wd14_tagger_loader_v3_node.py diff --git a/tests/test_wd14_tagger_runtime.py b/tests/segmentation/detection/test_wd14_tagger_runtime.py similarity index 100% rename from tests/test_wd14_tagger_runtime.py rename to tests/segmentation/detection/test_wd14_tagger_runtime.py diff --git a/tests/segmentation/segs/__init__.py b/tests/segmentation/segs/__init__.py new file mode 100644 index 0000000..8b9a266 --- /dev/null +++ b/tests/segmentation/segs/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own segmentation segs test behavior.""" diff --git a/tests/test_batch_segs_node.py b/tests/segmentation/segs/test_batch_segs_node.py similarity index 100% rename from tests/test_batch_segs_node.py rename to tests/segmentation/segs/test_batch_segs_node.py diff --git a/tests/test_batch_segs_v3_node.py b/tests/segmentation/segs/test_batch_segs_v3_node.py similarity index 100% rename from tests/test_batch_segs_v3_node.py rename to tests/segmentation/segs/test_batch_segs_v3_node.py diff --git a/tests/segmentation/segs/test_combined_segs_output.py b/tests/segmentation/segs/test_combined_segs_output.py new file mode 100644 index 0000000..8063601 --- /dev/null +++ b/tests/segmentation/segs/test_combined_segs_output.py @@ -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 + +"""Verify combined native-SEGS output construction.""" + +from __future__ import annotations + +from typing import cast + +import pytest +import torch + +from simple_syrup.domain.segs import BoundingBox, CropRegion, NativeSegs, Segment +from simple_syrup.services.segs_output_service import build_combined_segs_result + + +def test_empty_segs_build_empty_combined_outputs() -> None: + """Empty SEGS produce empty combined SEGS and a zero mask.""" + + result = build_combined_segs_result(_image(), ((8, 8), ()), crop_factor=1.0) + assert result.segs == ((8, 8), ()) + assert result.mask.shape == (1, 8, 8) + assert result.mask.dtype == torch.float32 + assert result.mask.sum().item() == 0.0 + + +def test_combined_result_unions_all_segment_masks() -> None: + """Combined SEGS contain one unioned segment for all source masks.""" + + image = torch.arange(1 * 8 * 8 * 3, dtype=torch.float32).reshape((1, 8, 8, 3)) + image = image / image.max() + first = _segment(CropRegion(1, 1, 3, 3), 0.4, torch.ones((2, 2))) + second = _segment( + CropRegion(4, 5, 7, 7), + 0.9, + [[2.0, 2.0, 2.0], [2.0, 2.0, 2.0]], + ) + result = build_combined_segs_result(image, _segs(first, second), crop_factor=1.0) + _header, segments = result.segs + assert result.mask.shape == (1, 8, 8) + assert result.mask.sum().item() == 10.0 + assert len(segments) == 1 + assert segments[0].crop_region == CropRegion(1, 1, 7, 7) + assert segments[0].bbox == BoundingBox(1, 1, 7, 7) + assert segments[0].label == "combined" + assert segments[0].confidence == 0.9 + assert cast(torch.Tensor, segments[0].cropped_mask).shape == (6, 6) + assert torch.equal( + cast(torch.Tensor, segments[0].cropped_image), image[:, 1:7, 1:7, :] + ) + + +def test_combined_result_applies_crop_factor_after_combining() -> None: + """Combined SEGS expand one unioned target for downstream detailers.""" + + image = torch.arange(1 * 8 * 8 * 3, dtype=torch.float32).reshape((1, 8, 8, 3)) + image = image / image.max() + first_mask = torch.zeros((2, 2), dtype=torch.float32) + first_mask[0, 0] = 1.0 + second_mask = torch.zeros((2, 2), dtype=torch.float32) + second_mask[1, 1] = 1.0 + first = Segment( + cropped_image=None, + cropped_mask=first_mask, + confidence=0.8, + crop_region=CropRegion(2, 2, 4, 4), + bbox=BoundingBox(2, 2, 3, 3), + label="face", + ) + second = Segment( + cropped_image=None, + cropped_mask=second_mask, + confidence=0.9, + crop_region=CropRegion(4, 4, 6, 6), + bbox=BoundingBox(5, 5, 6, 6), + label="face", + ) + result = build_combined_segs_result(image, _segs(first, second), crop_factor=2.0) + _header, segments = result.segs + assert len(segments) == 1 + assert segments[0].crop_region == CropRegion(0, 0, 8, 8) + assert segments[0].bbox == BoundingBox(2, 2, 6, 6) + assert cast(torch.Tensor, segments[0].cropped_mask).shape == (8, 8) + assert torch.equal(cast(torch.Tensor, segments[0].cropped_image), image) + + +def test_combined_mask_uses_max_for_overlaps() -> None: + """Overlapping source masks combine by max instead of addition.""" + + first = _segment(CropRegion(2, 2, 4, 4), 0.4, torch.full((2, 2), 0.7)) + second = _segment(CropRegion(2, 2, 4, 4), 0.8, torch.full((2, 2), 0.8)) + result = build_combined_segs_result(_image(), _segs(first, second), crop_factor=1.0) + assert result.mask[0, 2, 2].item() == pytest.approx(0.8) + assert result.mask.sum().item() == pytest.approx(3.2) + + +def test_combined_outputs_reject_mismatched_cropped_mask_shape() -> None: + """Invalid crop-local masks fail before producing misleading outputs.""" + + segment = _segment(CropRegion(1, 1, 3, 3), 1.0, torch.ones((3, 3))) + with pytest.raises(ValueError, match="cropped_mask must match"): + build_combined_segs_result(_image(), _segs(segment), crop_factor=1.0) + + +def _segs(*segments: Segment) -> NativeSegs: + """Return native SEGS for an 8x8 image.""" + + return (8, 8), tuple(segments) + + +def _segment( + crop_region: CropRegion, + confidence: float, + cropped_mask: object, +) -> Segment: + """Return one segment for combined-output tests.""" + + return Segment( + cropped_image=None, + cropped_mask=cropped_mask, + confidence=confidence, + crop_region=crop_region, + bbox=BoundingBox( + crop_region.left, + crop_region.top, + crop_region.right, + crop_region.bottom, + ), + label="face", + ) + + +def _image() -> torch.Tensor: + """Return a small single-image tensor.""" + + return torch.zeros((1, 8, 8, 3), dtype=torch.float32) diff --git a/tests/test_context_segs.py b/tests/segmentation/segs/test_context_segs.py similarity index 100% rename from tests/test_context_segs.py rename to tests/segmentation/segs/test_context_segs.py diff --git a/tests/test_detail_geometry.py b/tests/segmentation/segs/test_detail_geometry.py similarity index 100% rename from tests/test_detail_geometry.py rename to tests/segmentation/segs/test_detail_geometry.py diff --git a/tests/test_detail_masks.py b/tests/segmentation/segs/test_detail_masks.py similarity index 100% rename from tests/test_detail_masks.py rename to tests/segmentation/segs/test_detail_masks.py diff --git a/tests/test_detail_previews.py b/tests/segmentation/segs/test_detail_previews.py similarity index 100% rename from tests/test_detail_previews.py rename to tests/segmentation/segs/test_detail_previews.py diff --git a/tests/test_detail_resize.py b/tests/segmentation/segs/test_detail_resize.py similarity index 100% rename from tests/test_detail_resize.py rename to tests/segmentation/segs/test_detail_resize.py diff --git a/tests/test_detail_segs_as_regions_node.py b/tests/segmentation/segs/test_detail_segs_as_regions_node.py similarity index 100% rename from tests/test_detail_segs_as_regions_node.py rename to tests/segmentation/segs/test_detail_segs_as_regions_node.py diff --git a/tests/test_detail_segs_as_regions_service.py b/tests/segmentation/segs/test_detail_segs_as_regions_service.py similarity index 100% rename from tests/test_detail_segs_as_regions_service.py rename to tests/segmentation/segs/test_detail_segs_as_regions_service.py diff --git a/tests/test_detail_segs_by_scale_factor_node.py b/tests/segmentation/segs/test_detail_segs_by_scale_factor_node.py similarity index 100% rename from tests/test_detail_segs_by_scale_factor_node.py rename to tests/segmentation/segs/test_detail_segs_by_scale_factor_node.py diff --git a/tests/test_detail_segs_by_scale_factor_service.py b/tests/segmentation/segs/test_detail_segs_by_scale_factor_service.py similarity index 100% rename from tests/test_detail_segs_by_scale_factor_service.py rename to tests/segmentation/segs/test_detail_segs_by_scale_factor_service.py diff --git a/tests/test_detail_segs_by_scale_factor_tiled_diffusion_node.py b/tests/segmentation/segs/test_detail_segs_by_scale_factor_tiled_diffusion_node.py similarity index 100% rename from tests/test_detail_segs_by_scale_factor_tiled_diffusion_node.py rename to tests/segmentation/segs/test_detail_segs_by_scale_factor_tiled_diffusion_node.py diff --git a/tests/test_detail_segs_by_scale_factor_tiled_diffusion_service.py b/tests/segmentation/segs/test_detail_segs_by_scale_factor_tiled_diffusion_service.py similarity index 80% rename from tests/test_detail_segs_by_scale_factor_tiled_diffusion_service.py rename to tests/segmentation/segs/test_detail_segs_by_scale_factor_tiled_diffusion_service.py index 73be82f..10aa84e 100644 --- a/tests/test_detail_segs_by_scale_factor_tiled_diffusion_service.py +++ b/tests/segmentation/segs/test_detail_segs_by_scale_factor_tiled_diffusion_service.py @@ -17,7 +17,6 @@ from simple_syrup.runtime.detail_sampling import Latent from simple_syrup.services.detail_segs_by_scale_factor_tiled_diffusion_service import ( DetailSEGSByScaleFactorTiledDiffusionService, TiledDetailResizeBoundary, - TiledDetailSampler, TiledDetailSamplingBoundary, ) @@ -234,102 +233,6 @@ def test_tiled_noise_mask_feather_requests_single_clone_differential_diffusion() 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.""" - - tiled_sampling_service = _FakeTiledSamplingService() - latent = {"samples": torch.zeros((1, 4, 4, 4))} - preview_context = DetailPreviewContext( - image=_image(), - work_region=CropRegion(2, 2, 6, 6), - work_mask=torch.ones((4, 4)), - ) - - result = TiledDetailSampler( - tiled_sampling_service=tiled_sampling_service - ).sample_tiled( - diffusion_mode="mixture_of_diffusers", - model="model", - seed=123, - steps=4, - cfg=7.0, - sampler_name="euler", - scheduler="normal", - positive="positive", - negative="negative", - latent_image=latent, - denoise=0.5, - latent_tile_width=128, - latent_tile_height=80, - latent_tile_overlap=12, - latent_tile_batch_size=3, - preview_context=preview_context, - differential_diffusion=True, - ) - - assert result is latent - call = tiled_sampling_service.calls[0] - 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: - """Fake shared tiled diffusion sampling service.""" - - def __init__(self) -> None: - """Create empty call records.""" - - self.calls: list[dict[str, Any]] = [] - - def sample( - 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 = None, - differential_diffusion: bool = False, - ) -> Latent: - """Record tiled sampling arguments and return the latent unchanged.""" - - self.calls.append( - { - "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, - } - ) - return latent_image - - class _FakeTiledSampler: """Fake tiled sampling boundary for service tests.""" diff --git a/tests/test_mask_components.py b/tests/segmentation/segs/test_mask_components.py similarity index 100% rename from tests/test_mask_components.py rename to tests/segmentation/segs/test_mask_components.py diff --git a/tests/test_mask_ops.py b/tests/segmentation/segs/test_mask_ops.py similarity index 100% rename from tests/test_mask_ops.py rename to tests/segmentation/segs/test_mask_ops.py diff --git a/tests/test_mask_to_segs_node.py b/tests/segmentation/segs/test_mask_to_segs_node.py similarity index 100% rename from tests/test_mask_to_segs_node.py rename to tests/segmentation/segs/test_mask_to_segs_node.py diff --git a/tests/test_mask_to_segs_service.py b/tests/segmentation/segs/test_mask_to_segs_service.py similarity index 100% rename from tests/test_mask_to_segs_service.py rename to tests/segmentation/segs/test_mask_to_segs_service.py diff --git a/tests/test_prompt_segs_with_sam_compatibility.py b/tests/segmentation/segs/test_prompt_segs_with_sam_compatibility.py similarity index 98% rename from tests/test_prompt_segs_with_sam_compatibility.py rename to tests/segmentation/segs/test_prompt_segs_with_sam_compatibility.py index 6d2617e..45e91c1 100644 --- a/tests/test_prompt_segs_with_sam_compatibility.py +++ b/tests/segmentation/segs/test_prompt_segs_with_sam_compatibility.py @@ -11,9 +11,9 @@ from typing import cast import pytest import torch +from support.helpers import make_image_tensor from simple_syrup.masking.prompt_segs_with_sam_service import PromptSEGSWithSAMService -from test_helpers import make_image_tensor class PredictBoxesModel: diff --git a/tests/test_prompt_segs_with_sam_node.py b/tests/segmentation/segs/test_prompt_segs_with_sam_node.py similarity index 99% rename from tests/test_prompt_segs_with_sam_node.py rename to tests/segmentation/segs/test_prompt_segs_with_sam_node.py index 8d3847d..42963c1 100644 --- a/tests/test_prompt_segs_with_sam_node.py +++ b/tests/segmentation/segs/test_prompt_segs_with_sam_node.py @@ -9,6 +9,7 @@ from __future__ import annotations from typing import Any, cast import torch +from support.helpers import make_image_tensor from simple_syrup.domain.segs import ( SORT_ORDER_OPTIONS, @@ -18,7 +19,6 @@ from simple_syrup.domain.segs import ( ) from simple_syrup.nodes.prompt_segs_with_sam import PromptSEGSWithSAM from simple_syrup.services.segs_output_service import CombinedSegsResult -from test_helpers import make_image_tensor def test_prompt_segs_with_sam_node_contract_constants() -> None: diff --git a/tests/test_prompt_segs_with_sam_service.py b/tests/segmentation/segs/test_prompt_segs_with_sam_service.py similarity index 99% rename from tests/test_prompt_segs_with_sam_service.py rename to tests/segmentation/segs/test_prompt_segs_with_sam_service.py index 3cc528b..e73eab7 100644 --- a/tests/test_prompt_segs_with_sam_service.py +++ b/tests/segmentation/segs/test_prompt_segs_with_sam_service.py @@ -10,6 +10,7 @@ from typing import Protocol, cast import pytest import torch +from support.helpers import make_image_tensor from simple_syrup.domain.segs import BoundingBox, NativeSegs from simple_syrup.masking.prompt_segs_with_sam_service import ( @@ -17,7 +18,6 @@ from simple_syrup.masking.prompt_segs_with_sam_service import ( PromptSEGSWithSAMService, ) from simple_syrup.runtime.text_box_detector import TextBoxDetection -from test_helpers import make_image_tensor class RecordingDetector: diff --git a/tests/test_segs_domain.py b/tests/segmentation/segs/test_segs_domain.py similarity index 100% rename from tests/test_segs_domain.py rename to tests/segmentation/segs/test_segs_domain.py diff --git a/tests/test_segs_from_sam_output_node.py b/tests/segmentation/segs/test_segs_from_sam_output_node.py similarity index 100% rename from tests/test_segs_from_sam_output_node.py rename to tests/segmentation/segs/test_segs_from_sam_output_node.py diff --git a/tests/test_segs_from_sam_output_service.py b/tests/segmentation/segs/test_segs_from_sam_output_service.py similarity index 100% rename from tests/test_segs_from_sam_output_service.py rename to tests/segmentation/segs/test_segs_from_sam_output_service.py diff --git a/tests/test_segs_guided_tiled_diffusion_sampling_service.py b/tests/segmentation/segs/test_segs_guided_tiled_diffusion_sampling_service.py similarity index 100% rename from tests/test_segs_guided_tiled_diffusion_sampling_service.py rename to tests/segmentation/segs/test_segs_guided_tiled_diffusion_sampling_service.py diff --git a/tests/test_segs_output_service.py b/tests/segmentation/segs/test_segs_output_service.py similarity index 100% rename from tests/test_segs_output_service.py rename to tests/segmentation/segs/test_segs_output_service.py diff --git a/tests/test_tile_segs_domain.py b/tests/segmentation/segs/test_tile_segs_domain.py similarity index 100% rename from tests/test_tile_segs_domain.py rename to tests/segmentation/segs/test_tile_segs_domain.py diff --git a/tests/segmentation/segs/test_tiled_detail_sampler.py b/tests/segmentation/segs/test_tiled_detail_sampler.py new file mode 100644 index 0000000..672007d --- /dev/null +++ b/tests/segmentation/segs/test_tiled_detail_sampler.py @@ -0,0 +1,116 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Verify tiled detail sampling adapter delegation.""" + +from __future__ import annotations + +from typing import Any + +import torch + +from simple_syrup.domain.segs import CropRegion +from simple_syrup.runtime.detail_previews import DetailPreviewContext +from simple_syrup.runtime.detail_sampling import Latent +from simple_syrup.services.tiled_detail_sampler import TiledDetailSampler + + +def test_tiled_detail_sampler_delegates_to_shared_sampling_service() -> None: + """The detailer adapter uses the shared tiled diffusion dispatch service.""" + + tiled_sampling_service = _FakeTiledSamplingService() + latent = {"samples": torch.zeros((1, 4, 4, 4))} + preview_context = DetailPreviewContext( + image=_image(), + work_region=CropRegion(2, 2, 6, 6), + work_mask=torch.ones((4, 4)), + ) + result = TiledDetailSampler( + tiled_sampling_service=tiled_sampling_service + ).sample_tiled( + diffusion_mode="mixture_of_diffusers", + model="model", + seed=123, + steps=4, + cfg=7.0, + sampler_name="euler", + scheduler="normal", + positive="positive", + negative="negative", + latent_image=latent, + denoise=0.5, + latent_tile_width=128, + latent_tile_height=80, + latent_tile_overlap=12, + latent_tile_batch_size=3, + preview_context=preview_context, + differential_diffusion=True, + ) + assert result is latent + call = tiled_sampling_service.calls[0] + 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: + """Record calls to the shared tiled diffusion sampling service.""" + + def __init__(self) -> None: + """Create empty call records.""" + + self.calls: list[dict[str, Any]] = [] + + def sample( + 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 = None, + differential_diffusion: bool = False, + ) -> Latent: + """Record tiled sampling arguments and return the latent unchanged.""" + + self.calls.append( + { + "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, + } + ) + return latent_image + + +def _image() -> torch.Tensor: + """Return a small preview image.""" + + return torch.zeros((1, 8, 8, 3), dtype=torch.float32) diff --git a/tests/settings/__init__.py b/tests/settings/__init__.py new file mode 100644 index 0000000..653de30 --- /dev/null +++ b/tests/settings/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own settings test behavior.""" diff --git a/tests/test_quant_cache_routes.py b/tests/settings/test_quant_cache_routes.py similarity index 98% rename from tests/test_quant_cache_routes.py rename to tests/settings/test_quant_cache_routes.py index 7718203..5c01280 100644 --- a/tests/test_quant_cache_routes.py +++ b/tests/settings/test_quant_cache_routes.py @@ -13,7 +13,7 @@ from typing import cast from aiohttp import web -from simple_syrup.runtime.quant_cache_routes import ( +from simple_syrup.integration.quant_cache_routes import ( QUANT_CACHE_ROUTE, Handler, QuantCachePromptServerProtocol, diff --git a/tests/test_settings.py b/tests/settings/test_settings.py similarity index 100% rename from tests/test_settings.py rename to tests/settings/test_settings.py diff --git a/tests/test_settings_routes.py b/tests/settings/test_settings_routes.py similarity index 98% rename from tests/test_settings_routes.py rename to tests/settings/test_settings_routes.py index c921eea..97d40e4 100644 --- a/tests/test_settings_routes.py +++ b/tests/settings/test_settings_routes.py @@ -15,18 +15,18 @@ from typing import cast import pytest from aiohttp import web -import simple_syrup.runtime.settings_routes as settings_routes -from simple_syrup.runtime.settings import ( - ExternalLLMSettings, - SimpleSyrupSettings, -) -from simple_syrup.runtime.settings_repository import SimpleSyrupSettingsRepository -from simple_syrup.runtime.settings_routes import ( +import simple_syrup.integration.settings_routes as settings_routes +from simple_syrup.integration.settings_routes import ( SETTINGS_ROUTE, Handler, PromptServerProtocol, register_settings_routes, ) +from simple_syrup.runtime.settings import ( + ExternalLLMSettings, + SimpleSyrupSettings, +) +from simple_syrup.runtime.settings_repository import SimpleSyrupSettingsRepository class FakeRoutes: diff --git a/tests/support/__init__.py b/tests/support/__init__.py new file mode 100644 index 0000000..f7343ea --- /dev/null +++ b/tests/support/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own support test behavior.""" diff --git a/test_helpers.py b/tests/support/helpers.py similarity index 100% rename from test_helpers.py rename to tests/support/helpers.py diff --git a/tests/support/repository.py b/tests/support/repository.py new file mode 100644 index 0000000..d27708c --- /dev/null +++ b/tests/support/repository.py @@ -0,0 +1,12 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Expose stable repository paths to capability-owned tests.""" + +from __future__ import annotations + +from pathlib import Path + +REPOSITORY_ROOT = Path(__file__).resolve().parents[2] +CUSTOM_NODES_ROOT = REPOSITORY_ROOT.parent diff --git a/tests/test_attention_region_capture.py b/tests/test_attention_region_capture.py deleted file mode 100644 index a956451..0000000 --- a/tests/test_attention_region_capture.py +++ /dev/null @@ -1,1055 +0,0 @@ -# SimpleSyrup - workflow-focused ComfyUI extensions for image generation -# Copyright (C) 2026 Artificial Sweetener and contributors -# SPDX-License-Identifier: AGPL-3.0-or-later - -"""Test observation-only selected-token attention capture.""" - -from __future__ import annotations - -from concurrent.futures import ThreadPoolExecutor -from threading import Event -from typing import Any - -import pytest -import torch - -from simple_syrup.domain.attention_region_capture import ( - AttentionCapturePlan, - AttentionCaptureProfile, - AttentionRegionControls, - AttentionRegionRequest, - AttentionRegionRequestKind, -) -from simple_syrup.domain.attention_region_maps import ( - AttentionTokenCatalog, - AttentionTokenSpan, - OpenVocabularyContext, -) -from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily -from simple_syrup.runtime.attention_region_affinity import ( - ATTENTION_AFFINITY_CALCULATOR, - _select_relevant_head_maps, -) -from simple_syrup.runtime.attention_region_capture import AttentionRegionCaptureSession -from simple_syrup.runtime.attention_region_capture_backend import ( - ATTENTION_REGION_CAPTURE_BACKEND, - OptimizedAttentionCaptureOverride, -) -from simple_syrup.runtime.attention_region_contextual_spans import ( - ATTENTION_CONTEXTUAL_SPAN_SELECTOR, -) -from simple_syrup.runtime.attention_region_open_vocabulary import ( - OpenVocabularyKeyProjector, -) -from simple_syrup.runtime.attention_region_phrase_evidence import ( - ANIMA_PHRASE_EVIDENCE_SERVICE, - specific_attention_head_weights, -) -from simple_syrup.runtime.attention_region_self_completion import ( - ATTENTION_REGION_SELF_COMPLETION, - MAXIMUM_SELF_COMPLETION_ANCHORS, - SpatialSelfAttention, - _grid_anchor_indices, -) - - -def test_attention_controls_reject_invalid_instance_recall() -> None: - """Reject disconnected-instance recall outside its normalized range.""" - - with pytest.raises(ValueError, match="instance recall"): - AttentionRegionControls( - 0.0, - 1.0, - 0.15, - 0.25, - 0.0, - 1, - AttentionCaptureProfile.FAST, - instance_recall=1.01, - ) - - -def test_attention_controls_reject_invalid_geometry_recall() -> None: - """Reject connected-geometry recall outside its normalized range.""" - - with pytest.raises(ValueError, match="geometry recall"): - AttentionRegionControls( - 0.0, - 1.0, - 0.15, - 0.25, - 0.0, - 1, - AttentionCaptureProfile.FAST, - geometry_recall=-0.01, - ) - - -def test_capture_selects_positive_rows_and_exact_prompt_tokens() -> None: - """Capture one native concept without retaining a full attention matrix.""" - - session = _session(profile=AttentionCaptureProfile.EXHAUSTIVE) - query = torch.tensor( - [ - [[1.0, 0.0], [0.0, 1.0]], - [[-1.0, 0.0], [0.0, -1.0]], - ] - ) - key = torch.tensor( - [ - [[0.0, 0.0], [1.0, 0.0], [0.0, 1.0]], - [[0.0, 0.0], [-1.0, 0.0], [0.0, -1.0]], - ] - ) - - session.observe( - query, - key, - key, - 1, - _options(cond_or_uncond=[0, 1]), - skip_reshape=False, - ) - - maps = session.maps_for("search") - assert len(maps) == 1 - assert maps[0].label == "pink hair" - assert maps[0].values.device.type == "cpu" - assert maps[0].values[0].item() > maps[0].values[1].item() - - -def test_capture_retains_value_aware_phrase_evidence_beside_raw_attention() -> None: - """Favor the phrase token carrying stronger projected model contribution.""" - - session = _session( - profile=AttentionCaptureProfile.EXHAUSTIVE, - sequence_length=4, - token_indices=(1, 2), - ) - query = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]]) - key = torch.tensor([[[0.0, 0.0], [2.0, 0.0], [0.0, 2.0], [0.0, 0.0]]]) - value = torch.tensor([[[1.0, 1.0], [0.1, 0.1], [4.0, 4.0], [1.0, 1.0]]]) - - session.observe( - query, - key, - value, - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - attention_map = session.maps_for("search")[0] - assert torch.isclose(attention_map.values[0], attention_map.values[1]) - assert attention_map.concept_values is not None - assert attention_map.concept_values[1] > attention_map.concept_values[0] - assert attention_map.uniform_probability == 0.25 - - -def test_contextual_span_selector_admits_only_distinct_related_prompt_spans() -> None: - """Select a related conditioning span while rejecting an orthogonal concept.""" - - target = AttentionTokenSpan("pink hair", 1, (1,)) - related = AttentionTokenSpan("twintails", 1, (2,)) - unrelated = AttentionTokenSpan("smile", 1, (3,)) - key = torch.tensor([[[[0.0, 0.0], [1.0, 0.0], [0.9, 0.1], [0.0, 1.0]]]]) - - selections = ATTENTION_CONTEXTUAL_SPAN_SELECTOR.select( - key=key, - targets=(target,), - candidates=(target, related, unrelated), - ) - - assert selections[target].token_indices == (1, 2) - assert selections[target].token_weights[0] == 1.0 - assert 0.0 < selections[target].token_weights[1] < 1.0 - - -def test_contextual_span_selector_uses_object_head_instead_of_modifier_color() -> None: - """Relate an object phrase through its head without following a color token.""" - - target = AttentionTokenSpan("pink outfit", 1, (1, 2), (2,)) - related_part = AttentionTokenSpan("short skirt", 1, (3, 4), (4,)) - color_only = AttentionTokenSpan("pink petals", 1, (5,), (5,)) - key = torch.tensor( - [ - [ - [ - [0.0, 0.0], - [0.0, 1.0], - [1.0, 0.0], - [0.0, 1.0], - [0.9, 0.1], - [0.0, 1.0], - ] - ] - ] - ) - - selection = ATTENTION_CONTEXTUAL_SPAN_SELECTOR.select( - key=key, - targets=(target,), - candidates=(target, related_part, color_only), - )[target] - - assert selection.token_indices == (1, 2, 4) - assert selection.token_weights[0] < selection.token_weights[1] - - -def test_anima_compound_phrase_uses_specific_modifier_to_constrain_its_head() -> None: - """Keep a small compound concept local when its noun head is spatially broad.""" - - probability = torch.tensor( - [ - [ - [ - [0.8, 1.0, 1.0], - [0.8, 0.9, 0.9], - [0.8, 0.05, 0.7], - [0.8, 0.05, 0.7], - ] - ] - ] - ) - span = AttentionTokenSpan( - "blue butterfly ornaments", - 1, - (0, 1, 2), - (2,), - ) - - concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( - probability=probability, - span=span, - union_positions={0: 0, 1: 1, 2: 2}, - ) - - assert concept[0, 0] > 0.8 - assert concept[0, 1] > 0.7 - assert concept[0, 2] < concept[0, 0] * 0.5 - assert concept[0, 3] < concept[0, 0] * 0.5 - - -def test_anima_broad_modifier_does_not_erase_extended_head_geometry() -> None: - """Preserve a noun silhouette when its only modifier carries no locality.""" - - probability = torch.tensor( - [ - [ - [ - [0.6, 1.0], - [0.6, 0.8], - [0.6, 0.5], - [0.6, 0.3], - ] - ] - ] - ) - span = AttentionTokenSpan("pink hair", 1, (0, 1), (1,)) - - concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( - probability=probability, - span=span, - union_positions={0: 0, 1: 1}, - ) - - assert torch.allclose(concept[0], probability[0, 0, :, 1].to(torch.float16)) - - -def test_anima_single_token_concept_uses_noun_head_evidence_directly() -> None: - """Handle noun-only prompt segments without requiring modifier positions.""" - - probability = torch.tensor([[[[1.0], [0.8], [0.3], [0.0]]]]) - span = AttentionTokenSpan("twintails", 1, (0,), (0,)) - - concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( - probability=probability, - span=span, - union_positions={0: 0}, - ) - - assert torch.equal(concept[0], probability[0, 0, :, 0].to(torch.float16)) - - -def test_anima_related_prompt_head_recovers_a_disjoint_concept_part() -> None: - """Admit a related prompt noun at reduced strength without replacing the core.""" - - probability = torch.tensor( - [ - [ - [ - [0.5, 1.0, 0.0], - [0.5, 0.2, 0.0], - [0.5, 0.0, 0.9], - [0.5, 0.0, 0.8], - ] - ] - ] - ) - span = AttentionTokenSpan("pink hair", 1, (0, 1), (1,)) - - concept = ANIMA_PHRASE_EVIDENCE_SERVICE.derive( - probability=probability, - span=span, - union_positions={0: 0, 1: 1, 2: 2}, - contextual_token_indices=(0, 1, 2), - contextual_token_weights=(0.45, 1.0, 0.3), - ) - - assert concept[0, 0] == 1.0 - assert concept[0, 2] > 0.25 - assert concept[0, 3] > 0.2 - - -def test_concept_head_selection_rejects_a_spatially_disagreeing_head() -> None: - """Prefer concept heads that agree on the object while retaining their detail.""" - - probability = torch.tensor( - [ - [ - [[0.9], [0.8], [0.4], [0.0]], - [[0.8], [0.9], [0.0], [0.0]], - [[0.7], [0.8], [0.0], [0.0]], - [[0.0], [0.0], [0.9], [0.9]], - ] - ] - ) - - _weighted, selected = _select_relevant_head_maps( - probability, - torch.ones(1, 4, 1), - ) - - assert selected[0, 0, 0] > selected[0, 3, 0] - assert selected[0, 1, 0] > selected[0, 3, 0] - - -def test_related_tokens_use_heads_selected_by_the_exact_concept() -> None: - """Recover related detail without admitting its independently face-focused head.""" - - probability = torch.tensor( - [ - [ - [[0.9, 0.1], [0.8, 0.1], [0.4, 0.8], [0.0, 0.0]], - [[0.8, 0.1], [0.9, 0.1], [0.4, 0.7], [0.0, 0.0]], - [[0.7, 0.1], [0.8, 0.1], [0.3, 0.6], [0.0, 0.0]], - [[0.0, 0.0], [0.0, 0.0], [0.0, 0.1], [0.9, 0.9]], - ] - ] - ) - - _weighted, selected = _select_relevant_head_maps( - probability, - torch.ones(1, 4, 2), - exact_token_mask=torch.tensor([True, False]), - ) - - assert selected[0, 2, 1] > selected[0, 3, 1] - - -def test_concept_capture_enriches_related_span_without_changing_raw_attention() -> None: - """Recover self-grouped related evidence only in the derived concept channel.""" - - target = AttentionTokenSpan("pink hair", 1, (1,)) - related = AttentionTokenSpan("twintails", 1, (2,)) - unrelated = AttentionTokenSpan("smile", 1, (3,)) - session = _session( - profile=AttentionCaptureProfile.EXHAUSTIVE, - sequence_length=4, - catalog_spans=(target, related, unrelated), - ) - query = torch.tensor([[[1.0, 0.0], [0.0, 1.0]]]) - key = torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.6, 0.8], [0.0, -1.0]]]) - value = torch.ones_like(key) - - spatial = torch.tensor([[[1.0, 0.0], [1.0, 0.0]]]) - session.observe( - spatial, - spatial, - spatial, - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - session.observe( - query, - key, - value, - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - attention_map = session.maps_for("search")[0] - assert attention_map.concept_values is not None - raw_ratio = attention_map.values[1] / attention_map.values[0] - concept_ratio = attention_map.concept_values[1] / attention_map.concept_values[0] - assert concept_ratio > raw_ratio - - -def test_capture_pairs_spatial_self_attention_with_following_cross_call( - monkeypatch: Any, -) -> None: - """Pass same-layer self-attention only into derived concept evidence.""" - - session = _session(profile=AttentionCaptureProfile.EXHAUSTIVE) - observed_self_attention: list[SpatialSelfAttention | None] = [] - original = ATTENTION_AFFINITY_CALCULATOR.capture_spans - - def record_capture(*args: Any, **kwargs: Any) -> Any: - """Record staged self-attention before delegating to affinity capture.""" - - observed_self_attention.append(kwargs.get("self_attention")) - return original(*args, **kwargs) - - monkeypatch.setattr(ATTENTION_AFFINITY_CALCULATOR, "capture_spans", record_capture) - spatial = torch.tensor([[[1.0, 0.0], [0.9, 0.1], [0.8, 0.2], [0.0, 1.0]]]) - session.observe( - spatial, - spatial, - spatial, - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - session.observe( - spatial, - torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 1.0]]]), - torch.ones(1, 3, 2), - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - assert len(observed_self_attention) == 1 - assert isinstance(observed_self_attention[0], SpatialSelfAttention) - - -def test_anima_capture_does_not_apply_sdxl_self_attention_completion( - monkeypatch: Any, -) -> None: - """Keep Anima concept capture on its validated cross-attention evidence path.""" - - session = _session( - profile=AttentionCaptureProfile.EXHAUSTIVE, - sequence_length=512, - model_family=RegionalModelFamily.ANIMA, - source_aspect=0.75, - ) - observed_self_attention: list[SpatialSelfAttention | None] = [] - original = ATTENTION_AFFINITY_CALCULATOR.capture_spans - - def record_capture(*args: Any, **kwargs: Any) -> Any: - """Record completion input before delegating to affinity capture.""" - - observed_self_attention.append(kwargs.get("self_attention")) - return original(*args, **kwargs) - - monkeypatch.setattr(ATTENTION_AFFINITY_CALCULATOR, "capture_spans", record_capture) - spatial = torch.ones(1, 12, 2) - session.observe( - spatial, - spatial, - spatial, - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - session.observe( - spatial, - torch.ones(1, 512, 2), - torch.ones(1, 512, 2), - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - assert observed_self_attention == [None] - - -def test_anima_concept_evidence_preserves_contextualized_object_head() -> None: - """Let phrase modifiers refine Anima object evidence without erasing its extent.""" - - span = AttentionTokenSpan("pink hair", 1, (1, 2), (2,)) - session = _session( - profile=AttentionCaptureProfile.EXHAUSTIVE, - sequence_length=512, - model_family=RegionalModelFamily.ANIMA, - source_aspect=0.75, - catalog_spans=(span,), - ) - query = torch.tensor([[[2.0, 0.0], [0.0, 2.0], [1.5, 1.5]]]) - key = torch.zeros(1, 512, 2) - key[0, 1] = torch.tensor([1.0, 0.0]) - key[0, 2] = torch.tensor([0.0, 1.0]) - - session.observe( - query, - key, - torch.ones_like(key), - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - attention_map = session.maps_for("search")[0] - assert attention_map.concept_values is not None - assert attention_map.concept_values[1] > attention_map.concept_values[0] - assert attention_map.concept_values[1] > (attention_map.concept_values[2] * 0.75) - - -def test_anima_object_head_selection_rejects_diffuse_attention_heads() -> None: - """Prefer a spatially specific object head over broad interaction context.""" - - values = torch.tensor( - [ - [ - [0.1, 0.1, 0.9, 0.8], - [0.5, 0.5, 0.5, 0.5], - [0.4, 0.4, 0.4, 0.4], - [0.3, 0.3, 0.3, 0.3], - ] - ] - ) - - weights = specific_attention_head_weights(values) - - assert weights.shape == (1, 4) - assert weights[0, 0].item() == 1.0 - assert weights[0, 1:].count_nonzero().item() == 0 - - -def test_self_completion_drops_weak_grid_anchors() -> None: - """Keep spatial diversity without letting weak background cells steer completion.""" - - seed = torch.arange(9, dtype=torch.float32).reshape(1, 9) - - anchors = _grid_anchor_indices(seed, 3, 3) - - assert anchors.shape == (1, MAXIMUM_SELF_COMPLETION_ANCHORS) - assert set(anchors[0].tolist()) == {3, 4, 5, 6, 7, 8} - - -def test_self_completion_gates_related_recall_by_exact_object_grouping() -> None: - """Admit a related strand while rejecting an equally strong unrelated region.""" - - exact = torch.tensor([[1.0, 0.1, 0.05, 0.05]]) - related = torch.tensor([[1.0, 0.1, 0.9, 0.9]]) - query = torch.tensor([[[[1.0, 0.0], [0.0, 1.0], [1.0, 0.0], [0.0, 1.0]]]]) - self_attention = SpatialSelfAttention(query=query, key=query) - - completed = ATTENTION_REGION_SELF_COMPLETION.complete( - exact, - self_attention, - 2, - 2, - related_seed=related, - ) - - assert completed[0, 2] > completed[0, 3] - - -def test_sibling_consumers_share_one_concurrent_materialization( - monkeypatch: Any, -) -> None: - """Prevent parallel downstream nodes from consuming an emptied capture.""" - - session = _session(profile=AttentionCaptureProfile.EXHAUSTIVE) - session.observe( - torch.ones(1, 2, 2), - torch.ones(1, 3, 2), - torch.ones(1, 3, 2), - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - original = ATTENTION_AFFINITY_CALCULATOR.materialize - started = Event() - release = Event() - call_count = 0 - - def delayed_materialize(*args: Any, **kwargs: Any) -> Any: - """Hold the first materializer until a sibling is waiting.""" - - nonlocal call_count - call_count += 1 - started.set() - assert release.wait(timeout=2.0) - return original(*args, **kwargs) - - monkeypatch.setattr( - ATTENTION_AFFINITY_CALCULATOR, - "materialize", - delayed_materialize, - ) - with ThreadPoolExecutor(max_workers=2) as executor: - first = executor.submit(session.maps_for, "search") - assert started.wait(timeout=2.0) - second = executor.submit(session.maps_for, "search") - release.set() - results = (first.result(timeout=2.0), second.result(timeout=2.0)) - - assert call_count == 1 - assert all(len(result) == 1 for result in results) - assert results[0] is not results[1] - assert results[0][0].values.equal(results[1][0].values) - - -def test_capture_profile_subsamples_calls_without_changing_attention_output() -> None: - """Delegate every denoising call while retaining only fast-profile samples.""" - - session = _session(profile=AttentionCaptureProfile.FAST) - override = OptimizedAttentionCaptureOverride(session, None) - query = torch.ones(1, 2, 2) - key = torch.ones(1, 3, 2) - value = torch.arange(6, dtype=torch.float32).reshape(1, 3, 2) - - def original(*args: object, **kwargs: object) -> torch.Tensor: - """Return a sentinel output while accepting Comfy attention arguments.""" - - del args, kwargs - return torch.full((1, 2, 2), 7.0) - - options = tuple( - _options(cond_or_uncond=[0], block_index=index) for index in range(33) - ) - outputs = tuple( - override( - original, - query, - key, - value, - 1, - transformer_options=options[index], - ) - for index in range(33) - ) - - assert all(torch.equal(output, outputs[0]) for output in outputs) - assert outputs[0].eq(7.0).all().item() - assert len(session.maps_for("search")) == 2 - - -def test_fast_capture_rotates_sampled_layers_between_denoising_steps() -> None: - """Cover a different layer offset at each step without increasing stride cost.""" - - session = _session(profile=AttentionCaptureProfile.FAST) - query = torch.ones(1, 2, 2) - key = torch.ones(1, 3, 2) - value = torch.ones(1, 3, 2) - for block_index in range(33): - session.observe( - query, - key, - value, - 1, - _options(cond_or_uncond=[0], sigma=1.0, block_index=block_index), - skip_reshape=False, - ) - for block_index in range(2): - session.observe( - query, - key, - value, - 1, - _options(cond_or_uncond=[0], sigma=0.5, block_index=block_index), - skip_reshape=False, - ) - - assert len(session.maps_for("search")) == 3 - - -def test_fast_capture_subsamples_spatial_self_completion( - monkeypatch: Any, -) -> None: - """Retain sparse self-attention recall without paying for every fast sample.""" - - session = _session(profile=AttentionCaptureProfile.FAST) - observed: list[SpatialSelfAttention | None] = [] - original = ATTENTION_AFFINITY_CALCULATOR.capture_spans - - def capture_spans(*args: Any, **kwargs: Any) -> Any: - """Record completion inputs while preserving affinity behavior.""" - - observed.append(kwargs.get("self_attention")) - return original(*args, **kwargs) - - monkeypatch.setattr( - ATTENTION_AFFINITY_CALCULATOR, - "capture_spans", - capture_spans, - ) - for block_index in range(65): - options = _options(cond_or_uncond=[0], block_index=block_index) - session.observe( - torch.ones(1, 2, 2), - torch.ones(1, 2, 2), - torch.ones(1, 2, 2), - 1, - options, - skip_reshape=False, - ) - session.observe( - torch.ones(1, 2, 2), - torch.ones(1, 3, 2), - torch.ones(1, 3, 2), - 1, - options, - skip_reshape=False, - ) - - assert len(observed) == 3 - assert sum(isinstance(value, SpatialSelfAttention) for value in observed) == 1 - - -def test_coalesced_requests_share_one_unioned_native_affinity_pass( - monkeypatch: Any, -) -> None: - """Calculate shared prompt-token affinities once and route maps per request.""" - - controls = AttentionRegionControls( - 0.0, - 1.0, - 0.3, - 0.2, - 0.5, - 1, - AttentionCaptureProfile.EXHAUSTIVE, - ) - requests = ( - AttentionRegionRequest( - "all", - AttentionRegionRequestKind.ALL_PROMPT_SEGS, - (), - controls, - ), - AttentionRegionRequest( - "hair", - AttentionRegionRequestKind.CONCEPT_SEGS, - ("pink hair",), - controls, - ), - ) - plan = AttentionCapturePlan( - "sampler", - "sampler", - "model", - ("model", 0), - ("positive", 0), - requests, - "1girl, pink hair", - ("loader", 1), - ) - spans = ( - AttentionTokenSpan("1girl", 1, (1,)), - AttentionTokenSpan("pink hair", 1, (2,)), - ) - session = AttentionRegionCaptureSession( - plan=plan, - model_family=RegionalModelFamily.STANDARD_UNET, - token_catalog=AttentionTokenCatalog(4, spans, (0, 1, 2, 3)), - request_spans={"hair": (spans[1],), "all": spans}, - ) - original = ATTENTION_AFFINITY_CALCULATOR.capture_spans - calls: list[tuple[AttentionTokenSpan, ...]] = [] - - def record_capture(*args: Any, **kwargs: Any) -> Any: - """Record the unioned spans before delegating to real affinity math.""" - - captured_spans = kwargs.get("spans", args[3]) - calls.append(captured_spans) - return original(*args, **kwargs) - - monkeypatch.setattr(ATTENTION_AFFINITY_CALCULATOR, "capture_spans", record_capture) - session.observe( - torch.ones(1, 2, 2), - torch.ones(1, 4, 2), - torch.ones(1, 4, 2), - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - assert calls == [spans] - assert tuple(value.label for value in session.maps_for("hair")) == ("pink hair",) - assert tuple(value.label for value in session.maps_for("all")) == ( - "1girl", - "pink hair", - ) - - -def test_backend_clones_model_and_composes_existing_attention_override() -> None: - """Keep source options intact while preserving an upstream override.""" - - source = _patcher() - - def previous( - original: object, - *args: object, - **kwargs: object, - ) -> torch.Tensor: - """Return a sentinel from an admitted upstream override.""" - - del original, args, kwargs - return torch.tensor(3.0) - - source.model_options["transformer_options"]["optimized_attention_override"] = ( - previous - ) - - derived: Any = ATTENTION_REGION_CAPTURE_BACKEND.derive(source, _session()) - - assert derived.parent is source - assert ( - source.model_options["transformer_options"]["optimized_attention_override"] - is previous - ) - installed = derived.model_options["transformer_options"][ - "optimized_attention_override" - ] - assert isinstance(installed, OptimizedAttentionCaptureOverride) - assert ( - installed( - lambda *_args, **_kwargs: torch.tensor(1.0), - torch.ones(1, 1, 1), - torch.ones(1, 1, 1), - torch.ones(1, 1, 1), - 1, - ).item() - == 3.0 - ) - - -def test_open_vocabulary_query_projects_side_keys_without_changing_native_keys() -> ( - None -): - """Reuse SDXL spatial queries for an absent phrase without conditioning edits.""" - - controls = AttentionRegionControls( - 0.0, - 1.0, - 0.3, - 0.2, - 0.5, - 1, - AttentionCaptureProfile.EXHAUSTIVE, - ) - request = AttentionRegionRequest( - "search", - AttentionRegionRequestKind.CONCEPT_SEGS, - ("head",), - controls, - ) - plan = AttentionCapturePlan( - "sampler", - "sampler", - "model", - ("model", 0), - ("positive", 0), - (request,), - "1girl", - ("loader", 1), - ) - context = OpenVocabularyContext( - "head", - torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 0.0]]]), - (1,), - ) - session = AttentionRegionCaptureSession( - plan=plan, - model_family=RegionalModelFamily.STANDARD_UNET, - token_catalog=AttentionTokenCatalog( - 3, (AttentionTokenSpan("1girl", 1, (1,)),), (1, 2, 3) - ), - request_spans={"search": ()}, - open_vocabulary_contexts=(context,), - ) - projector = OpenVocabularyKeyProjector(torch.nn.Identity(), (context,), session) - native_context = torch.tensor([[[2.0, 0.0], [0.0, 2.0], [0.0, 0.0]]]) - - native_keys = projector(native_context) - session.observe( - torch.tensor([[[1.0, 0.0], [0.0, 1.0]]]), - native_keys, - native_context, - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - assert torch.equal(native_keys, native_context) - maps = session.maps_for("search") - assert len(maps) == 1 - assert maps[0].label == "head" - assert maps[0].values[0].item() > maps[0].values[1].item() - - -def test_open_vocabulary_query_projection_is_cached_per_device_and_dtype() -> None: - """Avoid repeating absent-query key projections at every denoising call.""" - - class CountingProjection(torch.nn.Module): - """Count native and side-projection calls while preserving values.""" - - calls: int - - def __init__(self) -> None: - """Initialize an unused projection counter.""" - - super().__init__() - self.calls = 0 - - def forward(self, value: torch.Tensor) -> torch.Tensor: - """Return the input after recording the projection.""" - - self.calls += 1 - return value - - context = OpenVocabularyContext( - "head", - torch.tensor([[[0.0, 0.0], [1.0, 0.0], [0.0, 0.0]]]), - (1,), - ) - projection = CountingProjection() - projector = OpenVocabularyKeyProjector(projection, (context,), _session()) - native_context = torch.ones(1, 3, 2) - - assert torch.equal(projector(native_context), native_context) - assert torch.equal(projector(native_context), native_context) - assert projection.calls == 3 - - -def test_sampled_denominator_bounds_unusually_strong_selected_keys() -> None: - """Keep fast-profile maps finite when denominator sampling misses the peak key.""" - - session = _session( - profile=AttentionCaptureProfile.FAST, - sequence_length=40, - ) - query = torch.tensor([[[1000.0, 0.0], [0.0, 1.0]]]) - key = torch.zeros(1, 40, 2) - key[0, 1, 0] = 1000.0 - - session.observe( - query, - key, - key, - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - maps = session.maps_for("search") - assert len(maps) == 1 - assert torch.isfinite(maps[0].values).all().item() - assert maps[0].values.max().item() <= 1.0 - - -def test_graph_source_aspect_orients_dit_geometry_without_transformer_metadata() -> ( - None -): - """Use the selected sampler's portrait input when a DiT omits shape metadata.""" - - session = _session(source_aspect=0.75) - - session.observe( - torch.ones(1, 12, 2), - torch.ones(1, 3, 2), - torch.ones(1, 3, 2), - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - attention_map = session.maps_for("search")[0] - assert (attention_map.spatial_height, attention_map.spatial_width) == (4, 3) - - -def test_anima_without_graph_or_runtime_geometry_fails_closed() -> None: - """Return no maps instead of guessing an unprojectable Anima grid orientation.""" - - session = _session( - sequence_length=512, - model_family=RegionalModelFamily.ANIMA, - ) - - session.observe( - torch.ones(1, 12, 2), - torch.ones(1, 512, 2), - torch.ones(1, 512, 2), - 1, - _options(cond_or_uncond=[0]), - skip_reshape=False, - ) - - assert session.maps_for("search") == () - assert "not graph-visible" in session.status_message - - -def _session( - profile: AttentionCaptureProfile = AttentionCaptureProfile.EXHAUSTIVE, - sequence_length: int = 3, - source_aspect: float | None = None, - model_family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, - token_indices: tuple[int, ...] = (1,), - catalog_spans: tuple[AttentionTokenSpan, ...] | None = None, -) -> AttentionRegionCaptureSession: - """Return one exact-native-query capture session.""" - - controls = AttentionRegionControls(0.0, 1.0, 0.3, 0.2, 0.5, 1, profile) - request = AttentionRegionRequest( - "search", - AttentionRegionRequestKind.CONCEPT_SEGS, - ("pink hair",), - controls, - ) - plan = AttentionCapturePlan( - "sampler", - "sampler", - "model", - ("model", 0), - ("positive", 0), - (request,), - "pink hair", - ("loader", 1), - source_aspect, - ) - default_span = AttentionTokenSpan("pink hair", 1, token_indices) - catalog = AttentionTokenCatalog( - sequence_length, - catalog_spans or (default_span,), - tuple(range(sequence_length)), - ) - return AttentionRegionCaptureSession( - plan=plan, - model_family=model_family, - token_catalog=catalog, - request_spans={"search": (catalog.spans[0],)}, - ) - - -def _options( - *, - cond_or_uncond: list[int], - sigma: float = 1.0, - block_index: int = 0, -) -> dict[str, object]: - """Return exact sampler metadata at the beginning of denoising.""" - - return { - "sample_sigmas": torch.tensor([1.0, 0.5, 0.0]), - "sigmas": torch.tensor([sigma]), - "cond_or_uncond": cond_or_uncond, - "block": ("middle", 0), - "block_index": block_index, - } - - -def _patcher() -> Any: - """Create a real CPU ModelPatcher with isolated transformer options.""" - - from comfy.model_patcher import ModelPatcher - - base_model = torch.nn.Module() - base_model.diffusion_model = torch.nn.Linear(1, 1) - device = torch.device("cpu") - return ModelPatcher(base_model, load_device=device, offload_device=device) diff --git a/tests/test_attention_region_rendering.py b/tests/test_attention_region_rendering.py deleted file mode 100644 index 7b604bb..0000000 --- a/tests/test_attention_region_rendering.py +++ /dev/null @@ -1,1298 +0,0 @@ -# SimpleSyrup - workflow-focused ComfyUI extensions for image generation -# Copyright (C) 2026 Artificial Sweetener and contributors -# SPDX-License-Identifier: AGPL-3.0-or-later - -"""Test attention-native temporal shaping and soft SEGS construction.""" - -from __future__ import annotations - -import torch - -from simple_syrup.domain.attention_region_capture import ( - AttentionCaptureProfile, - AttentionEvidenceMode, - AttentionRegionControls, -) -from simple_syrup.domain.attention_region_maps import CapturedAttentionMap -from simple_syrup.domain.regional_model_capabilities import RegionalModelFamily -from simple_syrup.services.attention_region_matte import ATTENTION_MATTE_SERVICE -from simple_syrup.services.attention_region_rendering import ( - ATTENTION_REGION_RENDERING_SERVICE, -) - - -def test_strength_and_consensus_shape_soft_regions_before_component_packaging() -> None: - """Reject transient weak pixels while preserving feathered accepted values.""" - - maps = ( - _map("hair", [0.1, 0.8, 0.2, 0.7], 0.2), - _map("hair", [0.1, 0.9, 0.1, 0.2], 0.5), - _map("hair", [0.1, 0.7, 0.1, 0.1], 0.8), - ) - image = torch.zeros(1, 2, 2, 3) - - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls(strength=0.5, consensus=0.66), - height=2, - width=2, - image=image, - ) - - assert len(segs[1]) == 1 - assert segs[1][0].label == "hair" - assert isinstance(segs[1][0].cropped_mask, torch.Tensor) - assert segs[1][0].cropped_mask.shape == (1, 1) - assert mask[0, 0, 1].item() == 1.0 - assert mask.count_nonzero().item() == 1 - - -def test_temporal_window_changes_region_without_morphology() -> None: - """Select early versus late attention evidence through normalized progress.""" - - maps = ( - _map("composition", [1.0, 1.0, 0.0, 0.0], 0.1), - _map("composition", [0.0, 0.0, 1.0, 1.0], 0.9), - ) - - _early_segs, early = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls(start=0.0, end=0.4, strength=0.2, consensus=0.0), - height=2, - width=2, - ) - _late_segs, late = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls(start=0.6, end=1.0, strength=0.2, consensus=0.0), - height=2, - width=2, - ) - - assert early[0, 0].sum().item() > early[0, 1].sum().item() - assert late[0, 1].sum().item() > late[0, 0].sum().item() - - -def test_raw_attention_preserves_diffuse_model_evidence() -> None: - """Keep low-amplitude positive attention visible in inspection mode.""" - - maps = ( - _map( - "outfit", - [0.26, 0.25, 0.27, 0.25], - 0.1, - baseline=0.25, - ), - _map( - "outfit", - [0.25, 0.80, 0.25, 0.25], - 0.6, - baseline=0.25, - ), - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.0, - consensus=0.0, - evidence_mode=AttentionEvidenceMode.RAW, - ), - height=2, - width=2, - ) - - assert mask.count_nonzero().item() == 4 - assert mask[0, 0, 1].item() == 1.0 - assert mask[0, 0, 0].item() > 0.0 - - -def test_concept_isolation_removes_uniform_attention_baseline() -> None: - """Do not promote near-uniform attention into concept support.""" - - maps = ( - _map( - "outfit", - [0.26, 0.25, 0.27, 0.25], - 0.1, - baseline=0.25, - ), - _map( - "outfit", - [0.25, 0.80, 0.25, 0.25], - 0.5, - baseline=0.25, - ), - _map( - "outfit", - [0.25, 0.75, 0.25, 0.25], - 0.7, - baseline=0.25, - ), - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.2, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=2, - width=2, - ) - - assert mask.count_nonzero().item() == 1 - assert mask[0, 0, 1].item() == 1.0 - - -def test_anima_concept_mode_uses_validated_probability_aggregation() -> None: - """Keep Anima on its spatially faithful cross-attention evidence policy.""" - - maps = ( - _map( - "cat", - [0.1, 0.8, 0.2, 0.7], - 0.2, - family=RegionalModelFamily.ANIMA, - concept_values=[0.8, 0.1, 0.7, 0.2], - baseline=0.1, - ), - _map( - "cat", - [0.1, 0.9, 0.1, 0.2], - 0.7, - family=RegionalModelFamily.ANIMA, - concept_values=[0.9, 0.1, 0.8, 0.1], - baseline=0.1, - ), - ) - concept_controls = _controls( - strength=0.2, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ) - raw_controls = _controls( - strength=0.2, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.RAW, - ) - - _concept_segs, concept_mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=concept_controls, - height=2, - width=2, - ) - _raw_segs, raw_mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=raw_controls, - height=2, - width=2, - ) - - assert not torch.equal(concept_mask, raw_mask) - assert concept_mask[0, 0, 0] > concept_mask[0, 0, 1] - assert raw_mask[0, 0, 1] > raw_mask[0, 0, 0] - - -def test_anima_concept_mode_recovers_connected_exact_token_geometry() -> None: - """Keep a raw-attention appendage attached to the semantic concept core.""" - - raw = [0.0] * 5 + [1.0, 1.0, 0.8, 0.7, 0.0] + [0.0] * 5 - concept = [0.0] * 5 + [1.0, 1.0, 0.0, 0.0, 0.0] + [0.0] * 5 - maps = tuple( - _map( - "cat", - raw, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=concept, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.5, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=3, - width=5, - ) - - assert mask[0, 1, :4].count_nonzero().item() == 4 - assert mask.count_nonzero().item() == 4 - - -def test_anima_concept_mode_rejects_detached_exact_token_noise() -> None: - """Exclude raw-attention components that do not touch the semantic core.""" - - maps = tuple( - _map( - "cat", - [1.0, 1.0, 0.0, 0.0, 0.9], - progress, - family=RegionalModelFamily.ANIMA, - concept_values=[1.0, 1.0, 0.0, 0.0, 0.0], - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.5, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=1, - width=5, - ) - - assert mask[0, 0, :2].count_nonzero().item() == 2 - assert mask[0, 0, 4].item() == 0.0 - - -def test_anima_concept_mode_preserves_complete_expansive_anchored_geometry() -> None: - """Keep the complete attached object while dropping its weak global bridge.""" - - maps = tuple( - _map( - "mage staff", - [1.0, 1.0, 0.8, 0.8, 0.8, 0.2, 0.2, 0.2, 0.2, 0.2], - progress, - family=RegionalModelFamily.ANIMA, - concept_values=[1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0], - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=1, - width=10, - ) - - assert mask[0, 0, :5].count_nonzero().item() == 5 - assert mask[0, 0, 5:].count_nonzero().item() == 0 - - -def test_anima_concept_mode_favors_recall_when_expansion_is_ambiguous() -> None: - """Keep attached geometry when no tighter support extends beyond the core.""" - - maps = tuple( - _map( - "close subject", - [1.0, 1.0, 0.8, 0.8, 0.8], - progress, - family=RegionalModelFamily.ANIMA, - concept_values=[1.0, 1.0, 0.0, 0.0, 0.0], - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=1, - width=5, - ) - - assert mask.count_nonzero().item() == 5 - - -def test_anima_concept_mode_rejects_weak_broad_field_around_compact_peak() -> None: - """Keep a compact semantic peak without absorbing its weak connected field.""" - - raw = [0.0] * 100 - for row in range(3, 7): - for column in range(10): - raw[row * 10 + column] = 0.2 - raw[44] = 1.0 - raw[45] = 1.0 - concept = [0.0] * 100 - concept[44] = 1.0 - concept[45] = 1.0 - maps = tuple( - _map( - "compact feature", - raw, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=concept, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=0.85, - ), - height=10, - width=10, - ) - - assert mask.count_nonzero().item() == 2 - - -def test_anima_concept_mode_rejects_global_reference_geometry_for_local_core() -> None: - """Do not expand localized evidence through a near-global exact-token field.""" - - raw = [0.3] * 100 - concept = [0.0] * 100 - for row in range(3, 7): - for column in range(10): - raw[row * 10 + column] = 0.4 - for column in range(3, 8): - concept[4 * 10 + column] = 1.0 - maps = tuple( - _map( - "localized feature", - raw, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=concept, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=0.85, - ), - height=10, - width=10, - ) - - assert mask.count_nonzero().item() == 5 - - -def test_anima_concept_mode_prefers_concentrated_geometry_for_moderate_core() -> None: - """Use stricter geometry consensus when a moderate core anchors a broad field.""" - - raw = [0.0] * 100 - for row in range(3, 7): - for column in range(10): - raw[row * 10 + column] = 0.2 - concept = [0.0] * 100 - for column in range(10): - raw[4 * 10 + column] = 1.0 - concept[4 * 10 + column] = 1.0 - maps = tuple( - _map( - "moderate compact feature", - raw, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=concept, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=0.85, - ), - height=10, - width=10, - ) - - assert mask.count_nonzero().item() == 10 - - -def test_anima_concept_mode_tightens_compact_geometry_below_frame_threshold() -> None: - """Use the stable core when a compact weak field occupies under 15% of a frame.""" - - raw = [0.0] * 400 - for row in range(6, 13): - for column in range(6, 13): - raw[row * 20 + column] = 0.2 - concept = [0.0] * 400 - for row in range(8, 11): - for column in range(8, 11): - raw[row * 20 + column] = 1.0 - concept[row * 20 + column] = 1.0 - maps = tuple( - _map( - "compact sub-frame feature", - raw, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=concept, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=0.85, - ), - height=20, - width=20, - ) - - assert mask.count_nonzero().item() == 9 - - -def test_full_geometry_recall_preserves_weak_broad_field_around_compact_peak() -> None: - """Let an explicit maximum-recall choice bypass adaptive compactness.""" - - raw = [0.0] * 100 - for row in range(3, 7): - for column in range(10): - raw[row * 10 + column] = 0.2 - raw[44] = 1.0 - raw[45] = 1.0 - concept = [0.0] * 100 - concept[44] = 1.0 - concept[45] = 1.0 - maps = tuple( - _map( - "compact feature", - raw, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=concept, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=1.0, - ), - height=10, - width=10, - ) - - assert mask.count_nonzero().item() == 40 - - -def test_anima_concept_mode_tightens_peak_dominated_semantic_field() -> None: - """Tighten broad support when strength belongs mainly to a small core.""" - - values = [0.05] * 100 - for row in range(3, 7): - for column in range(10): - values[row * 10 + column] = 0.3 - values[44] = 1.0 - values[45] = 1.0 - maps = tuple( - _map( - "compact semantic feature", - values, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=values, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=0.85, - ), - height=10, - width=10, - ) - - assert mask.count_nonzero().item() == 2 - - -def test_anima_concept_mode_preserves_broad_high_strength_semantic_region() -> None: - """Keep a genuinely broad concept whose support remains strong across its area.""" - - values = [0.05] * 100 - for row in range(2, 8): - for column in range(10): - values[row * 10 + column] = 0.8 - values[44] = 1.0 - values[45] = 1.0 - maps = tuple( - _map( - "broad semantic region", - values, - progress, - family=RegionalModelFamily.ANIMA, - concept_values=values, - ) - for progress in (0.2, 0.7) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=0.85, - ), - height=10, - width=10, - ) - - assert mask.count_nonzero().item() == 60 - - -def test_geometry_recall_controls_faint_connected_exact_token_extent() -> None: - """Let users trade faint attached geometry for a tighter semantic core.""" - - maps = tuple( - _map( - "cat tail", - [ - 1.0, - 1.0, - 0.6, - 0.6, - 0.4, - 0.4, - 0.2, - 0.2, - 0.0, - 0.0, - 0.0, - 0.0, - 0.0, - 0.0, - 0.0, - 0.0, - ], - progress, - family=RegionalModelFamily.ANIMA, - concept_values=[1.0, 1.0] + [0.0] * 14, - ) - for progress in (0.2, 0.7) - ) - - _recalled_segs, recalled = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=1.0, - ), - height=1, - width=16, - ) - _tight_segs, tight = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - geometry_recall=0.0, - ), - height=1, - width=16, - ) - - assert recalled.count_nonzero().item() == 8 - assert tight.count_nonzero().item() == 6 - - -def test_anima_concept_mode_rejects_below_baseline_late_residue() -> None: - """Do not normalize negligible late-step Anima residue into full support.""" - - maps = tuple( - _map( - "cuffs", - [0.01, 0.01, 0.01, 0.01], - progress, - family=RegionalModelFamily.ANIMA, - concept_values=[0.0010, 0.0012, 0.0011, 0.0010], - baseline=0.01, - ) - for progress in (0.75, 0.9) - ) - - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=2, - width=2, - ) - - assert segs[1] == () - assert mask.count_nonzero().item() == 0 - - -def test_anima_concept_mode_removes_a_broad_contextual_field() -> None: - """Keep local lift without treating a broadly elevated phrase as full-frame.""" - - maps = tuple( - _map( - "swept bangs", - [0.1, 0.1, 0.1, 0.1], - progress, - family=RegionalModelFamily.ANIMA, - concept_values=[0.11, 0.11, 0.11, 0.14], - baseline=0.1, - ) - for progress in (0.2, 0.6) - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.25, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=2, - width=2, - ) - - assert mask.count_nonzero().item() == 1 - assert mask[0, 1, 1].item() == 1.0 - - -def test_anima_concept_mode_prefers_resolved_mid_pass_evidence() -> None: - """Prevent an endpoint observation from tying resolved middle evidence.""" - - maps = ( - _map( - "boots", - [0.1, 0.1, 0.1, 0.1], - 0.0, - family=RegionalModelFamily.ANIMA, - concept_values=[0.9, 0.1, 0.1, 0.1], - baseline=0.1, - layer="shared", - ), - _map( - "boots", - [0.1, 0.1, 0.1, 0.1], - 0.5, - family=RegionalModelFamily.ANIMA, - concept_values=[0.1, 0.1, 0.1, 0.9], - baseline=0.1, - layer="shared", - ), - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.4, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=2, - width=2, - ) - - assert mask.count_nonzero().item() == 1 - assert mask[0, 1, 1].item() == 1.0 - - -def test_concept_isolation_prefers_repeated_support_over_transient_peak() -> None: - """Suppress one strong flash when stable observations localize elsewhere.""" - - maps = ( - _map( - "hair", - [0.25, 0.25, 0.25, 0.95], - 0.1, - baseline=0.25, - layer="early", - ), - _map( - "hair", - [0.25, 0.78, 0.25, 0.25], - 0.4, - baseline=0.25, - layer="middle-a", - ), - _map( - "hair", - [0.25, 0.82, 0.25, 0.25], - 0.6, - baseline=0.25, - layer="middle-b", - ), - _map( - "hair", - [0.25, 0.76, 0.25, 0.25], - 0.8, - baseline=0.25, - layer="late", - ), - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.2, - consensus=0.4, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=2, - width=2, - ) - - assert mask.count_nonzero().item() == 1 - assert mask[0, 0, 1].item() == 1.0 - - -def test_concept_isolation_preserves_repeated_fine_layer_structure() -> None: - """Keep layer-local fine detail without admitting a one-step transient.""" - - maps = tuple( - _map( - "hair", - [0.25, 0.85, 0.25, 0.25], - progress, - baseline=0.25, - layer="coarse", - ) - for progress in (0.2, 0.4, 0.6, 0.8) - ) + ( - _map( - "hair", - [0.25, 0.85, 0.52, 0.48], - 0.4, - baseline=0.25, - layer="fine", - ), - _map( - "hair", - [0.25, 0.85, 0.55, 0.25], - 0.7, - baseline=0.25, - layer="fine", - ), - ) - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.5, - feather=0, - evidence_mode=AttentionEvidenceMode.CONCEPT, - ), - height=2, - width=2, - ) - - assert mask[0, 1, 0].item() > 0.0 - assert mask[0, 1, 1].item() == 0.0 - - -def test_empty_maps_return_correctly_sized_no_op_outputs() -> None: - """Return empty SEGS and a zero mask for unsupported graph/model paths.""" - - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(), - controls=_controls(), - height=5, - width=7, - ) - - assert segs == ((5, 7), ()) - assert mask.shape == (1, 5, 7) - assert mask.count_nonzero().item() == 0 - - -def test_disconnected_attention_islands_become_separate_instances() -> None: - """Expose retained disconnected objects as independently controllable SEGS.""" - - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("outdoors", [1.0, 0.0, 1.0], 0.5),), - controls=_controls(strength=0.5, consensus=0.0), - height=1, - width=3, - ) - - assert len(segs[1]) == 2 - assert tuple(segment.label for segment in segs[1]) == ("outdoors", "outdoors") - assert mask.count_nonzero().item() == 2 - - -def test_concept_isolation_rejects_weak_disconnected_context() -> None: - """Drop a weak contextual island while retaining raw inspection evidence.""" - - maps = (_map("bangs", [1.0, 0.9, 0.0, 0.3, 0.3], 0.5),) - concept_segs, _concept_mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls(strength=0.15, consensus=0.0), - height=1, - width=5, - ) - raw_segs, _raw_mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.15, - consensus=0.0, - evidence_mode=AttentionEvidenceMode.RAW, - ), - height=1, - width=5, - ) - - assert len(concept_segs[1]) == 1 - assert concept_segs[1][0].bbox == (0, 0, 2, 1) - assert len(raw_segs[1]) == 2 - - -def test_concept_isolation_retains_multiple_confident_regions() -> None: - """Keep plural concept instances when each has substantial evidence.""" - - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("cuffs", [1.0, 0.9, 0.0, 0.7, 0.65], 0.5),), - controls=_controls(strength=0.15, consensus=0.0), - height=1, - width=5, - ) - - assert len(segs[1]) == 2 - assert mask.count_nonzero().item() == 4 - - -def test_concept_isolation_retains_sparse_instances_by_peak_evidence() -> None: - """Keep small repeated instances without rewarding a larger region for area.""" - - values = [0.6] * 16 + [0.0, 1.0, 0.0, 0.7] - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("falling petals", values, 0.5),), - controls=_controls(strength=0.15, consensus=0.0), - height=1, - width=len(values), - ) - - assert len(segs[1]) == 3 - assert mask[0, 0, 17].item() > 0.0 - assert mask[0, 0, 19].item() > 0.0 - - -def test_instance_recall_can_restrict_sparse_results_to_the_strongest_peak() -> None: - """Let users remove weaker disconnected instances without an area heuristic.""" - - values = [0.6] * 16 + [0.0, 1.0, 0.0, 0.7] - segs, _mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("falling petals", values, 0.5),), - controls=_controls( - strength=0.15, - consensus=0.0, - instance_recall=0.2, - ), - height=1, - width=len(values), - ) - - assert len(segs[1]) == 1 - assert segs[1][0].bbox == (17, 0, 18, 1) - - -def test_concept_isolation_does_not_prefer_a_tiny_sharp_island() -> None: - """Keep broad supported evidence when a disconnected pixel peaks higher.""" - - values = [0.4] * 25 + [0.0, 1.0] - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("hair", values, 0.5),), - controls=_controls(strength=0.15, consensus=0.0), - height=1, - width=len(values), - ) - - assert any(segment.bbox == (0, 0, 25, 1) for segment in segs[1]) - assert mask[0, 0, :25].count_nonzero().item() == 25 - - -def test_concept_isolation_rejects_pockmarks_around_a_compact_dominant_region() -> None: - """Keep one compact body when much smaller disconnected peaks surround it.""" - - values = torch.zeros(11, 11) - values[3:8, 3:8] = 0.6 - values[0:2, 0:2] = 0.8 - values[0, 10] = 1.0 - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("torso", values.flatten().tolist(), 0.5),), - controls=_controls(strength=0.15, consensus=0.0, feather=0), - height=11, - width=11, - ) - - assert len(segs[1]) == 1 - assert segs[1][0].bbox == (3, 3, 8, 8) - assert mask.count_nonzero().item() == 25 - - -def test_full_instance_recall_preserves_pockmarks_around_a_compact_region() -> None: - """Honor an explicit request to retain every supported disconnected instance.""" - - values = torch.zeros(11, 11) - values[3:8, 3:8] = 0.6 - values[0:2, 0:2] = 0.8 - values[0, 10] = 1.0 - segs, _mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("torso", values.flatten().tolist(), 0.5),), - controls=_controls( - strength=0.15, - consensus=0.0, - feather=0, - instance_recall=1.0, - ), - height=11, - width=11, - ) - - assert len(segs[1]) == 3 - - -def test_split_sensitivity_preserves_the_complete_concept_union() -> None: - """Change instance separation without deleting moderate concept support.""" - - maps = (_map("cat", [1.0, 0.4, 0.4, 0.4, 1.0], 0.5),) - _unsplit_segs, unsplit = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls(strength=0.3, consensus=0.0, split=0.0), - height=1, - width=5, - ) - split_segs, split = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls(strength=0.3, consensus=0.0, split=1.0), - height=1, - width=5, - ) - - assert torch.equal(split, unsplit) - assert split.count_nonzero().item() == 5 - assert len(split_segs[1]) == 2 - - -def test_cohesive_support_retains_a_moderate_body_around_a_strong_core() -> None: - """Keep the complete attended silhouette instead of its sparse core pixels.""" - - values = [ - 0.0, - 0.2, - 0.2, - 0.2, - 0.0, - 0.2, - 0.4, - 0.6, - 0.4, - 0.2, - 0.2, - 0.6, - 1.0, - 0.6, - 0.2, - 0.2, - 0.4, - 0.6, - 0.4, - 0.2, - 0.0, - 0.2, - 0.2, - 0.2, - 0.0, - ] - - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("hair", values, 0.5),), - controls=_controls(strength=0.15, consensus=0.0, split=0.75), - height=5, - width=5, - ) - - assert mask.count_nonzero().item() == 21 - assert mask[0, 2, 2].item() == 1.0 - assert mask[0, 0, 1].item() > 0.0 - - -def test_minimum_region_size_removes_each_pockmark_independently() -> None: - """Discard a tiny island without rejecting or merging the valid component.""" - - values = [1.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0] - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("hair", values, 0.5),), - controls=_controls(strength=0.5, consensus=0.0, minimum_size=2), - height=3, - width=3, - ) - - assert len(segs[1]) == 1 - assert segs[1][0].bbox == (0, 0, 2, 1) - assert mask.count_nonzero().item() == 2 - - -def test_keep_only_and_combine_apply_per_concept_without_changing_union() -> None: - """Retain top instances and optionally package their union as one SEG.""" - - maps = (_map("cat", [1.0, 1.0, 0.0, 0.8, 0.8], 0.5),) - separate, separate_mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.5, - consensus=0.0, - keep_only=2, - combine=False, - ), - height=1, - width=5, - ) - combined, combined_mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.5, - consensus=0.0, - keep_only=2, - combine=True, - ), - height=1, - width=5, - ) - - assert len(separate[1]) == 2 - assert len(combined[1]) == 1 - assert torch.equal(separate_mask, combined_mask) - - -def test_keep_largest_groups_a_nearby_detached_concept_fragment() -> None: - """Treat qualifying nearby support as one instance before top-N ranking.""" - - values = torch.zeros((9, 9), dtype=torch.float32) - values[1, 1:8] = 1.0 - values[7, 1:8] = 1.0 - values[1:8, 1] = 1.0 - values[4:6, 4:6] = 0.8 - - segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("cat", values.flatten().tolist(), 0.5),), - controls=_controls( - strength=0.5, - consensus=0.0, - keep_only=1, - feather=0, - ), - height=9, - width=9, - ) - - assert len(segs[1]) == 1 - assert mask[0, 4:6, 4:6].count_nonzero().item() == 4 - - -def test_matte_solidity_flattens_interior_and_edge_feather_softens_boundary() -> None: - """Preserve raw alpha at zero and make a solid feathered matte at one.""" - - maps = (_map("dress", [0.0, 0.6, 1.0, 0.7, 0.0], 0.5),) - _raw_segs, raw = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls(strength=0.2, consensus=0.0, solidity=0.0), - height=1, - width=5, - ) - _solid_segs, solid = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=maps, - controls=_controls( - strength=0.2, - consensus=0.0, - solidity=1.0, - feather=1, - ), - height=1, - width=5, - ) - - assert raw[0, 0, 1].item() != raw[0, 0, 2].item() - assert solid[0, 0, 2].item() == 1.0 - assert 0.0 < solid[0, 0, 0].item() < 1.0 - - -def test_full_matte_solidity_fills_only_enclosed_holes() -> None: - """Fill an interior attention gap without filling exterior-connected space.""" - - ring = [1.0, 1.0, 1.0, 1.0, 0.0, 1.0, 1.0, 1.0, 1.0] - _segs, mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("hair", ring, 0.5),), - controls=_controls( - strength=0.5, - consensus=0.0, - solidity=1.0, - feather=0, - ), - height=3, - width=3, - ) - - assert torch.equal(mask, torch.ones_like(mask)) - - -def test_full_matte_solidity_preserves_a_narrow_exterior_connected_channel() -> None: - """Flatten alpha without inventing support inside an exterior channel.""" - - support = torch.zeros(7, 7, dtype=torch.bool) - support[1:6, 1:6] = True - support[1:4, 3] = False - - matte = ATTENTION_MATTE_SERVICE.shape( - alpha=support.float(), - support=support, - solidity=1.0, - edge_feather=0, - ) - - assert matte[3, 3].item() == 0.0 - assert matte[0].count_nonzero().item() == 0 - assert matte[:, 0].count_nonzero().item() == 0 - - -def test_full_matte_solidity_preserves_a_winding_exterior_channel() -> None: - """Keep exterior-connected exclusions regardless of their shape or width.""" - - support = torch.ones(100, 100, dtype=torch.bool) - support[55:, 25:75] = False - support[35:55, 49:52] = False - support[42:45, 42:52] = False - - matte = ATTENTION_MATTE_SERVICE.shape( - alpha=support.float(), - support=support, - solidity=1.0, - edge_feather=0, - ) - - assert matte[38, 50].item() == 0.0 - assert matte[43, 44].item() == 0.0 - assert matte[80, 50].item() == 0.0 - - -def test_full_matte_solidity_never_erases_accepted_thin_support() -> None: - """Keep every accepted pixel when topology cleanup adds cohesive support.""" - - support = torch.zeros(100, 100, dtype=torch.bool) - support[10:90, 50] = True - - matte = ATTENTION_MATTE_SERVICE.shape( - alpha=support.float(), - support=support, - solidity=1.0, - edge_feather=0, - ) - - assert torch.all(matte[support] == 1.0) - - -def test_highest_confidence_keeps_stronger_component() -> None: - """Rank components from pre-normalized evidence rather than normalized maxima.""" - - segs, _mask = ATTENTION_REGION_RENDERING_SERVICE.render( - maps=(_map("cat", [1.0, 0.0, 0.55], 0.5),), - controls=AttentionRegionControls( - 0.0, - 1.0, - 0.4, - 0.0, - 0.0, - 1, - AttentionCaptureProfile.BALANCED, - keep_only=1, - keep_by="highest confidence", - ), - height=1, - width=3, - ) - - assert len(segs[1]) == 1 - assert segs[1][0].bbox == (0, 0, 1, 1) - - -def _map( - label: str, - values: list[float], - progress: float, - *, - baseline: float = 0.0, - layer: str = "layer", - family: RegionalModelFamily = RegionalModelFamily.STANDARD_UNET, - concept_values: list[float] | None = None, -) -> CapturedAttentionMap: - """Create one compact square attention observation.""" - - return CapturedAttentionMap( - label, - torch.tensor(values, dtype=torch.float16), - progress, - layer, - uniform_probability=baseline, - model_family=family, - concept_values=( - torch.tensor(concept_values, dtype=torch.float16) - if concept_values is not None - else None - ), - ) - - -def _controls( - *, - start: float = 0.0, - end: float = 1.0, - strength: float = 0.35, - consensus: float = 0.25, - minimum_size: int = 1, - keep_only: int = 0, - combine: bool = False, - solidity: float = 0.0, - feather: int = 8, - split: float = 0.0, - instance_recall: float = 0.65, - geometry_recall: float = 0.85, - evidence_mode: AttentionEvidenceMode = AttentionEvidenceMode.CONCEPT, -) -> AttentionRegionControls: - """Return representative balanced rendering controls.""" - - return AttentionRegionControls( - start, - end, - strength, - consensus, - split, - minimum_size, - AttentionCaptureProfile.BALANCED, - instance_recall=instance_recall, - geometry_recall=geometry_recall, - keep_only=keep_only, - combine_segs=combine, - matte_solidity=solidity, - edge_feather=feather, - evidence_mode=evidence_mode, - ) diff --git a/tests/test_regional_attention_selection.py b/tests/test_regional_attention_selection.py deleted file mode 100644 index 73482c2..0000000 --- a/tests/test_regional_attention_selection.py +++ /dev/null @@ -1,338 +0,0 @@ -# SimpleSyrup - workflow-focused ComfyUI extensions for image generation -# Copyright (C) 2026 Artificial Sweetener and contributors -# SPDX-License-Identifier: AGPL-3.0-or-later - -"""Prove exact current-sigma regional conditioning entry selection.""" - -from __future__ import annotations - -from uuid import UUID, uuid4 - -import pytest -import torch - -from simple_syrup.domain.conditioning_schedule import ConditioningScheduleRange -from simple_syrup.domain.processed_regional_attention import ( - ProcessedRegionalAttentionBranch, - ProcessedRegionalAttentionContext, - ProcessedRegionalAttentionEntry, - ProcessedRegionalAttentionPlan, -) -from simple_syrup.domain.regional_attention_selection import ( - REGIONAL_ATTENTION_SELECTION_SERVICE, -) -from simple_syrup.domain.regional_lora_plan import RegionalLoraPlan -from simple_syrup.domain.regional_mask_bank import RegionalMaskBank -from simple_syrup.runtime.regional_attention_batching import ( - REGIONAL_ATTENTION_BATCHING_SERVICE, -) - - -@pytest.mark.parametrize( - ("sigma", "active_indices"), - [ - (100.0, (0,)), - (75.0, (0, 2)), - (50.0, (0, 1, 2, 3)), - (25.0, (1, 2)), - (0.0, (1,)), - ], -) -def test_selection_preserves_all_active_entries_for_positive_and_negative_banks( - sigma: float, - active_indices: tuple[int, ...], -) -> None: - """Keep inclusive adjacent, overlap, strength, identity, and source order.""" - - plan = _scheduled_plan() - positive_base = _active_base_entry(plan.positive, sigma=sigma) - negative_base = _active_base_entry(plan.negative, sigma=sigma) - - chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0, 1], - conditioning_uuids=[positive_base.uuid, negative_base.uuid], - sigma=sigma, - ) - - expected_positive = tuple( - plan.positive.regional_contexts[0].entries[index] for index in active_indices - ) - expected_negative = tuple( - plan.negative.regional_contexts[0].entries[index] for index in active_indices - ) - assert chunks[0].base_entry is positive_base - assert chunks[1].base_entry is negative_base - assert chunks[0].regional_entries == (expected_positive,) - assert chunks[1].regional_entries == (expected_negative,) - assert [entry.strength for entry in expected_positive] == [ - (0.1, 0.2, 0.3, 0.4)[index] for index in active_indices - ] - assert [entry.strength for entry in expected_negative] == [ - (0.1, 0.2, 0.3, 0.4)[index] for index in active_indices - ] - - -def test_selection_uses_exact_uuid_when_adjacent_base_entries_are_both_active() -> None: - """Resolve equality-boundary chunks by UUID instead of branch occurrence.""" - - plan = _scheduled_plan() - positive_entries = plan.positive.base_context.entries - - chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0, 0], - conditioning_uuids=[positive_entries[1].uuid, positive_entries[0].uuid], - sigma=50.0, - ) - - assert [chunk.base_entry for chunk in chunks] == [ - positive_entries[1], - positive_entries[0], - ] - - -def test_selection_retains_authored_no_active_region_without_base_fallback() -> None: - """Represent a schedule gap distinctly from an absent regional context.""" - - plan = _gap_plan() - uuids = [ - plan.positive.base_context.entries[0].uuid, - plan.negative.base_context.entries[0].uuid, - ] - - chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0, 1], - conditioning_uuids=uuids, - sigma=50.0, - ) - aligned = REGIONAL_ATTENTION_BATCHING_SERVICE.align( - plan, - base_context=torch.cat( - ( - plan.positive.base_context.entries[0].cross_attention.repeat(2, 1, 1), - plan.negative.base_context.entries[0].cross_attention.repeat(2, 1, 1), - ) - ), - cond_or_uncond=[0, 1], - conditioning_uuids=uuids, - sigma=50.0, - latent_batch_size=2, - ) - - assert chunks[0].regional_entries == ((),) - assert chunks[1].regional_entries == ((),) - assert len(aligned.regions[0].entries) == 1 - assert aligned.regions[0].entries[0].strengths == (0.0, 0.0, 0.0, 0.0) - - -def test_selection_keeps_absent_region_as_explicit_base_fallback() -> None: - """Distinguish no authored regional context from an inactive authored one.""" - - plan = _gap_plan(negative_region=False) - uuids = [ - plan.positive.base_context.entries[0].uuid, - plan.negative.base_context.entries[0].uuid, - ] - - chunks = REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0, 1], - conditioning_uuids=uuids, - sigma=50.0, - ) - aligned = REGIONAL_ATTENTION_BATCHING_SERVICE.align( - plan, - base_context=torch.cat( - ( - plan.positive.base_context.entries[0].cross_attention, - plan.negative.base_context.entries[0].cross_attention, - ) - ), - cond_or_uncond=[0, 1], - conditioning_uuids=uuids, - sigma=50.0, - latent_batch_size=1, - ) - - assert chunks[0].regional_entries == ((),) - assert chunks[1].regional_entries == (None,) - assert aligned.regions[0].entries[0].strengths == (0.0, 1.0) - - -def test_selection_rejects_unknown_or_inactive_base_uuid() -> None: - """Fail closed when supplied Comfy identity cannot be active in its branch.""" - - plan = _scheduled_plan() - inactive = plan.positive.base_context.entries[1] - with pytest.raises(ValueError, match="inactive at sigma"): - REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0], - conditioning_uuids=[inactive.uuid], - sigma=75.0, - ) - with pytest.raises(ValueError, match="does not identify exactly one"): - REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[1], - conditioning_uuids=[uuid4()], - sigma=75.0, - ) - equal_but_distinct_uuid = UUID(str(plan.negative.base_context.entries[0].uuid)) - with pytest.raises(ValueError, match="does not identify exactly one"): - REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[1], - conditioning_uuids=[equal_but_distinct_uuid], - sigma=75.0, - ) - - -@pytest.mark.parametrize( - ("uuids", "sigma", "error", "message"), - [ - ([], 50.0, ValueError, "equal lengths"), - ([object()], 50.0, TypeError, "Comfy UUID"), - (None, 50.0, TypeError, "list or tuple"), - ("valid", True, TypeError, "sigma must be a real"), - ("valid", float("nan"), ValueError, "sigma must be finite"), - ], -) -def test_selection_rejects_malformed_identity_or_sigma_state( - uuids: object, - sigma: object, - error: type[Exception], - message: str, -) -> None: - """Reject ambiguous model-call identity and timestep metadata.""" - - plan = _scheduled_plan() - conditioning_uuids = ( - [plan.positive.base_context.entries[0].uuid] if uuids == "valid" else uuids - ) - with pytest.raises(error, match=message): - REGIONAL_ATTENTION_SELECTION_SERVICE.select_chunks( - plan, - cond_or_uncond=[0], - conditioning_uuids=conditioning_uuids, - sigma=sigma, # type: ignore[arg-type] - ) - - -def _scheduled_plan() -> ProcessedRegionalAttentionPlan: - """Build symmetric branches with adjacent and overlapping entry schedules.""" - - return _plan( - positive_base=( - _entry(0, 1.0, start=100.0, end=50.0), - _entry(1, 2.0, start=50.0, end=0.0), - ), - negative_base=( - _entry(0, -1.0, start=100.0, end=50.0), - _entry(1, -2.0, start=50.0, end=0.0), - ), - positive_region=_scheduled_region(positive=True), - negative_region=_scheduled_region(positive=False), - ) - - -def _gap_plan(*, negative_region: bool = True) -> ProcessedRegionalAttentionPlan: - """Build branches whose authored regional entries share an inactive gap.""" - - gap_entries = ( - _entry(0, 2.0, start=100.0, end=75.0), - _entry(1, 3.0, start=25.0, end=0.0), - ) - return _plan( - positive_base=(_entry(0, 1.0),), - negative_base=(_entry(0, -1.0),), - positive_region=gap_entries, - negative_region=( - tuple( - _entry( - entry.entry_index, - -2.0, - start=entry.schedule.timestep_start, - end=entry.schedule.timestep_end, - ) - for entry in gap_entries - ) - if negative_region - else None - ), - ) - - -def _plan( - *, - positive_base: tuple[ProcessedRegionalAttentionEntry, ...], - negative_base: tuple[ProcessedRegionalAttentionEntry, ...], - positive_region: tuple[ProcessedRegionalAttentionEntry, ...], - negative_region: tuple[ProcessedRegionalAttentionEntry, ...] | None, -) -> ProcessedRegionalAttentionPlan: - """Build one-region processed plan from explicit entry banks.""" - - return ProcessedRegionalAttentionPlan( - ProcessedRegionalAttentionBranch( - ProcessedRegionalAttentionContext(0, None, positive_base), - (ProcessedRegionalAttentionContext(1, 0, positive_region),), - ), - ProcessedRegionalAttentionBranch( - ProcessedRegionalAttentionContext(0, None, negative_base), - ( - () - if negative_region is None - else (ProcessedRegionalAttentionContext(1, 0, negative_region),) - ), - ), - RegionalMaskBank(torch.ones((1, 2, 2)), torch.ones((1, 2, 2)), 2, 2), - RegionalLoraPlan(()), - ) - - -def _scheduled_region( - *, - positive: bool, -) -> tuple[ProcessedRegionalAttentionEntry, ...]: - """Build one recognizable scheduled region for either branch.""" - - sign = 1.0 if positive else -1.0 - return ( - _entry(0, 10.0 * sign, start=100.0, end=50.0, strength=0.1), - _entry(1, 20.0 * sign, start=50.0, end=0.0, strength=0.2), - _entry(2, 30.0 * sign, start=75.0, end=25.0, strength=0.3), - _entry(3, 40.0 * sign, start=50.0, end=50.0, strength=0.4), - ) - - -def _entry( - entry_index: int, - value: float, - *, - start: float | None = None, - end: float | None = None, - strength: float = 1.0, -) -> ProcessedRegionalAttentionEntry: - """Build one processed entry with explicit converted schedule boundaries.""" - - return ProcessedRegionalAttentionEntry( - entry_index, - uuid4(), - ConditioningScheduleRange(None, None, start, end), - torch.full((1, 2, 3), value), - strength, - ) - - -def _active_base_entry( - branch: ProcessedRegionalAttentionBranch, - *, - sigma: float, -) -> ProcessedRegionalAttentionEntry: - """Return the single active base entry away from the adjacent boundary.""" - - if sigma >= 50.0: - return branch.base_context.entries[0] - return branch.base_context.entries[1] diff --git a/tests/test_regional_multidiffusion_sampling.py b/tests/test_regional_multidiffusion_sampling.py deleted file mode 100644 index 486180d..0000000 --- a/tests/test_regional_multidiffusion_sampling.py +++ /dev/null @@ -1,975 +0,0 @@ -# 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 regional MultiDiffusion sampling runtime.""" - -from __future__ import annotations - -from types import SimpleNamespace -from typing import Any, cast - -import pytest -import torch - -from simple_syrup.domain.regional_detailing import LatentBox, LatentRegion -from simple_syrup.domain.segs import CropRegion -from simple_syrup.runtime import ( - regional_multidiffusion_sampling, - sampling_samplers, - sampling_schedulers, -) -from simple_syrup.runtime.detail_previews import DetailPreviewContext - -comfy_sample = regional_multidiffusion_sampling._comfy_sample() -comfy_utils = regional_multidiffusion_sampling._comfy_utils() -latent_preview = regional_multidiffusion_sampling._latent_preview() - - -def _preview_context() -> DetailPreviewContext: - """Return a minimal regional detail preview context.""" - - return DetailPreviewContext( - image=torch.ones((1, 8, 8, 3), dtype=torch.float32), - work_region=CropRegion(2, 2, 6, 6), - work_mask=torch.ones((8, 8), dtype=torch.float32), - sampled_region=CropRegion(0, 0, 8, 8), - ) - - -def test_sampling_callback_uses_generic_preview_without_detail_context( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Standalone regional runtime calls keep generic latent previews.""" - - monkeypatch.setattr( - latent_preview, - "prepare_callback", - lambda _model, _steps: "generic callback", - ) - - assert ( - regional_multidiffusion_sampling._sampling_callback(FakeModel(), 4, None) - == "generic callback" - ) - - -def test_sampling_callback_uses_detail_preview_with_context( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Detailer regional sampling uses the shared detail preview callback.""" - - context = _preview_context() - calls: dict[str, object] = {} - - def fake_prepare_detail_preview_callback( - model: FakeModel, - steps: int, - preview_context: DetailPreviewContext, - ) -> str: - """Record detail preview callback preparation.""" - - calls["model"] = model - calls["steps"] = steps - calls["preview_context"] = preview_context - return "detail callback" - - monkeypatch.setattr( - regional_multidiffusion_sampling, - "prepare_detail_preview_callback", - fake_prepare_detail_preview_callback, - ) - - model = FakeModel() - assert ( - regional_multidiffusion_sampling._sampling_callback(model, 4, context) - == "detail callback" - ) - assert calls == {"model": model, "steps": 4, "preview_context": context} - - -class FakeModel: - """Provide the ModelPatcher methods used by the runtime.""" - - def __init__( - self, - model_options: dict[str, Any] | None = None, - parent: FakeModel | None = None, - ) -> None: - """Create a fake model patcher.""" - - self.load_device = torch.device("cpu") - 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.""" - - 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.""" - - 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.""" - - assert name == "model_sampling" - return self.model_sampling - - -class FakeSampler: - """Represent a resolved sampler in tests.""" - - def sample(self, *args: object, **kwargs: object) -> object: - """Provide ComfyUI's sampler protocol.""" - - del args, kwargs - return None - - -def test_clone_model_installs_regional_calc_cond_batch_wrapper() -> None: - """The runtime clones the model and installs a regional wrapper.""" - - 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),), - ) - ) - - assert wrapped_model is not model - assert isinstance( - wrapped_model.calc_wrapper, - regional_multidiffusion_sampling.RegionalMultiDiffusionCalcCondBatch, - ) - 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.""" - - model = FakeModel({"sampler_calc_cond_batch_function": object()}) - - with pytest.raises(ValueError, match="Existing sampler_calc_cond_batch_function"): - 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),), - ) - - -def test_existing_calc_wrapper_is_composed_without_recursion() -> None: - """Fallback and regional calls delegate to the previous calc wrapper.""" - - calls: list[dict[str, Any]] = [] - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Return condition-specific constants and record options.""" - - calls.append(args) - assert args["model_options"].get("sampler_calc_cond_batch_function") is existing - x = cast(torch.Tensor, args["input"]) - conds = cast(list[object], args["conds"]) - value = 10.0 if _condition_name(conds[0]) == "global" else 20.0 - return [torch.ones_like(x) * value, torch.ones_like(x) * 2.0] - - model = FakeModel( - { - "sampler_calc_cond_batch_function": existing, - "model_function_wrapper": object(), - } - ) - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - model, - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=(_region(0, 0, 4, 4, "regional"),), - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert len(calls) == 2 - assert wrapped_model.model_options["model_function_wrapper"] is not None - assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 20.0) - assert torch.allclose(output[1], torch.ones((1, 1, 4, 4)) * 2.0) - - -def test_shape_mismatch_delegates_to_original_calc_path() -> None: - """Unexpected model input spatial shapes are delegated unchanged.""" - - calls = 0 - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Record fallback calls.""" - - nonlocal calls - calls += 1 - x = cast(torch.Tensor, args["input"]) - return [x + 5.0, x + 1.0] - - model = FakeModel({"sampler_calc_cond_batch_function": existing}) - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - model, - latent_width=8, - latent_height=8, - latent_ndim=4, - regions=(_region(0, 0, 4, 4, "regional", latent_width=8, latent_height=8),), - ) - ) - x = torch.zeros((1, 1, 4, 4)) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": x, - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert calls == 1 - assert torch.allclose(output[0], x + 5.0) - - -def test_region_crops_bchw_final_spatial_axes() -> None: - """Regional calls crop only final height and width axes.""" - - calls: list[torch.Tensor] = [] - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Record regional input crops.""" - - x = cast(torch.Tensor, args["input"]) - calls.append(x) - conds = cast(list[object], args["conds"]) - value = 1.0 if _condition_name(conds[0]) == "global" else 3.0 - return [torch.ones_like(x) * value, torch.zeros_like(x)] - - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=8, - latent_height=4, - latent_ndim=4, - regions=(_region(0, 0, 8, 4, "regional", latent_width=8),), - ) - ) - x = torch.arange(32, dtype=torch.float32).reshape((1, 1, 4, 8)) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": x, - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert calls[1].shape == (1, 1, 4, 8) - assert torch.equal(calls[1], x) - assert torch.allclose(output[0], torch.ones((1, 1, 4, 8)) * 3.0) - - -def test_region_crops_singleton_depth_5d_final_spatial_axes() -> None: - """Anima-style regions crop only final height and width axes.""" - - calls: list[torch.Tensor] = [] - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Record regional input crops.""" - - x = cast(torch.Tensor, args["input"]) - calls.append(x) - conds = cast(list[object], args["conds"]) - value = 1.0 if _condition_name(conds[0]) == "global" else 4.0 - return [torch.ones_like(x) * value, torch.zeros_like(x)] - - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=8, - latent_height=4, - latent_ndim=5, - regions=(_region(0, 0, 8, 4, "regional", latent_width=8),), - ) - ) - x = ( - torch.arange(128, dtype=torch.float32) - .reshape((1, 16, 1, 4, 2)) - .repeat(1, 1, 1, 1, 4) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": x, - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert calls[1].shape == (1, 16, 1, 4, 8) - assert torch.equal(calls[1], x) - assert output[0].shape == x.shape - - -def test_overlapping_regions_normalize_by_accumulated_weight() -> None: - """Overlapping regions are averaged before blending over fallback.""" - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Return different constants for each region positive.""" - - x = cast(torch.Tensor, args["input"]) - cond = _condition_name(cast(list[object], args["conds"])[0]) - values = {"global": 0.0, "first": 2.0, "second": 6.0} - return [torch.ones_like(x) * values[cond], torch.zeros_like(x)] - - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=6, - latent_height=4, - latent_ndim=4, - regions=( - _region(0, 0, 4, 4, "first", latent_width=6), - _region(2, 0, 4, 4, "second", latent_width=6), - ), - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": torch.zeros((1, 1, 4, 6)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert torch.allclose(output[0][:, :, :, :2], torch.ones((1, 1, 4, 2)) * 2.0) - assert torch.allclose(output[0][:, :, :, 2:4], torch.ones((1, 1, 4, 2)) * 4.0) - assert torch.allclose(output[0][:, :, :, 4:], torch.ones((1, 1, 4, 2)) * 6.0) - - -def test_partial_mask_blends_region_over_fallback() -> None: - """Feathered masks blend region predictions with fallback predictions.""" - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Return fallback or regional constants.""" - - x = cast(torch.Tensor, args["input"]) - cond = _condition_name(cast(list[object], args["conds"])[0]) - value = 10.0 if cond == "global" else 20.0 - return [torch.ones_like(x) * value, torch.zeros_like(x)] - - mask = torch.ones((4, 4)) * 0.25 - region = LatentRegion( - index=0, - label="soft", - latent_box=LatentBox(0, 0, 4, 4), - latent_mask=mask, - positive=_raw_conditioning("regional"), - ) - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=(region,), - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 12.5) - - -def test_raw_region_conditioning_is_converted_before_calc_cond_batch( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Raw CONDITIONING entries are converted before Comfy calc-cond-batch calls.""" - - calls: list[list[list[dict[str, Any]] | None]] = [] - - def calc_cond_batch( - _model: object, - conds: list[list[dict[str, Any]] | None], - x_in: torch.Tensor, - _timestep: torch.Tensor, - _model_options: dict[str, Any], - ) -> list[torch.Tensor]: - """Return constants while asserting sampler-ready conditioning shape.""" - - calls.append(conds) - first_cond = conds[0] - assert first_cond is not None - assert isinstance(first_cond[0], dict) - assert "model_conds" in first_cond[0] - value = 5.0 if "cross_attn" in first_cond[0] else 1.0 - return [torch.ones_like(x_in) * value, torch.zeros_like(x_in)] - - fake_samplers = SimpleNamespace( - calc_cond_batch=calc_cond_batch, - resolve_areas_and_cond_masks_multidim=lambda *_args: None, - calculate_start_end_timesteps=lambda *_args: None, - ) - monkeypatch.setattr( - regional_multidiffusion_sampling, - "_comfy_samplers", - lambda: fake_samplers, - ) - raw_region_positive = [[torch.ones((1, 1, 1)), {}]] - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel(), - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=(_region(0, 0, 4, 4, raw_region_positive),), - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": [ - [{"model_conds": {}, "uuid": object()}], - [{"model_conds": {}, "uuid": object()}], - ], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert len(calls) == 2 - assert "cross_attn" in cast(list[dict[str, Any]], calls[1][0])[0] - assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 5.0) - - -def test_global_prompt_weight_blends_full_region_with_global_prediction() -> None: - """Covered pixels keep the configured global positive prediction share.""" - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Return fallback or regional constants.""" - - x = cast(torch.Tensor, args["input"]) - cond = _condition_name(cast(list[object], args["conds"])[0]) - value = 10.0 if cond == "global" else 20.0 - return [torch.ones_like(x) * value, torch.zeros_like(x)] - - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=(_region(0, 0, 4, 4, "regional"),), - global_prompt_weight=0.25, - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 17.5) - - -def test_global_prompt_weight_keeps_partial_mask_coverage() -> None: - """Soft masks scale regional influence before global/regional weighting.""" - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Return fallback or regional constants.""" - - x = cast(torch.Tensor, args["input"]) - cond = _condition_name(cast(list[object], args["conds"])[0]) - value = 10.0 if cond == "global" else 20.0 - return [torch.ones_like(x) * value, torch.zeros_like(x)] - - region = LatentRegion( - index=0, - label="soft", - latent_box=LatentBox(0, 0, 4, 4), - latent_mask=torch.ones((4, 4)) * 0.5, - positive=_raw_conditioning("regional"), - ) - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=(region,), - global_prompt_weight=0.25, - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 13.75) - - -def test_overlapping_regions_normalize_before_global_prompt_weight_blend() -> None: - """Overlaps average region predictions before applying global weight.""" - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Return different constants for each condition.""" - - x = cast(torch.Tensor, args["input"]) - cond = _condition_name(cast(list[object], args["conds"])[0]) - values = {"global": 10.0, "first": 20.0, "second": 40.0} - return [torch.ones_like(x) * values[cond], torch.zeros_like(x)] - - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=( - _region(0, 0, 4, 4, "first"), - _region(0, 0, 4, 4, "second"), - ), - global_prompt_weight=0.25, - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 25.0) - - -def test_negative_conditioning_is_reused_for_region_unconditional_path() -> None: - """Regional calls keep the original negative conditioning.""" - - regional_conds: list[list[object]] = [] - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Record all conditioning lists.""" - - x = cast(torch.Tensor, args["input"]) - conds = cast(list[object], args["conds"]) - regional_conds.append(conds) - return [torch.ones_like(x), torch.ones_like(x) * 2.0] - - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=(_region(0, 0, 4, 4, "regional"),), - ) - ) - - wrapped_model.calc_wrapper( - { - "conds": ["global", "negative"], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert [_condition_name(item) for item in regional_conds[1]] == [ - "regional", - "negative", - ] - - -def test_cfg_one_none_uncond_still_returns_two_entries() -> None: - """Comfy's CFG=1 optimization keeps the output list shape stable.""" - - def existing(args: dict[str, Any]) -> list[torch.Tensor]: - """Return one tensor per cond slot, even when uncond is None.""" - - x = cast(torch.Tensor, args["input"]) - conds = cast(list[object], args["conds"]) - value = 3.0 if _condition_name(conds[0]) == "regional" else 1.0 - return [torch.ones_like(x) * value, torch.zeros_like(x)] - - wrapped_model, _summary = ( - regional_multidiffusion_sampling.clone_model_with_regional_multidiffusion( - FakeModel({"sampler_calc_cond_batch_function": existing}), - latent_width=4, - latent_height=4, - latent_ndim=4, - regions=(_region(0, 0, 4, 4, "regional"),), - ) - ) - - output = wrapped_model.calc_wrapper( - { - "conds": ["global", None], - "input": torch.zeros((1, 1, 4, 4)), - "sigma": torch.tensor([1.0]), - "model": wrapped_model, - "model_options": wrapped_model.model_options, - } - ) - - assert len(output) == 2 - assert torch.allclose(output[0], torch.ones((1, 1, 4, 4)) * 3.0) - assert torch.allclose(output[1], torch.zeros((1, 1, 4, 4))) - - -def test_sample_rejects_unipc_before_sampler_resolution( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """UniPC sampler names fail before ComfyUI sampler lookup.""" - - def fail_resolve_sampler(_sampler_name: str) -> FakeSampler: - """Fail if sampler resolution is reached.""" - - raise AssertionError("UniPC rejection should happen before sampler resolution") - - monkeypatch.setattr(sampling_samplers, "resolve_sampler", fail_resolve_sampler) - - with pytest.raises(ValueError, match="not compatible with UniPC"): - regional_multidiffusion_sampling.sample_regional_multidiffusion( - model=FakeModel(), - seed=1, - steps=1, - cfg=1.0, - sampler_name="uni_pc", - scheduler="normal", - positive=[], - negative=[], - latent_image={"samples": torch.zeros((1, 4, 4, 4))}, - regions=(_region(0, 0, 4, 4, "regional"),), - denoise=1.0, - global_prompt_weight=0.0, - ) - - -def test_sample_rejects_unsupported_conditioning() -> None: - """Regional and ControlNet conditioning fail closed.""" - - with pytest.raises(ValueError, match="regional conditioning or ControlNet"): - regional_multidiffusion_sampling.sample_regional_multidiffusion( - model=FakeModel(), - seed=1, - steps=1, - cfg=1.0, - sampler_name="euler", - scheduler="normal", - positive=[{"area": (4, 4, 0, 0)}], - negative=[], - latent_image={"samples": torch.zeros((1, 4, 4, 4))}, - regions=(_region(0, 0, 4, 4, "regional"),), - denoise=1.0, - global_prompt_weight=0.0, - ) - - -def test_sample_delegates_to_comfy_sampling_with_cloned_wrapped_model( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Sampling mirrors KSampler flow while using a wrapped model clone.""" - - calls: dict[str, Any] = {} - model = FakeModel() - sampler = FakeSampler() - latent_samples = torch.zeros((1, 4, 4, 8), dtype=torch.float32) - fixed_noise = torch.ones_like(latent_samples) - fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) - sampled = torch.full_like(latent_samples, 0.25) - latent_image: dict[str, Any] = { - "samples": latent_samples, - "downscale_ratio_spacial": 2, - "kept": "value", - } - - 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=None: fixed_noise, - ) - monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) - monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", False) - - 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: object, - disable_pbar: bool, - seed: int, - ) -> torch.Tensor: - """Record sample_custom arguments.""" - - calls["sample_custom"] = { - "model": received_model, - "noise": noise, - "cfg": cfg, - "sampler": received_sampler, - "sigmas": sigmas, - "positive": positive, - "negative": negative, - "latent_image": latent_image, - "noise_mask": noise_mask, - "callback": callback, - "disable_pbar": disable_pbar, - "seed": seed, - } - return sampled - - monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) - - output = regional_multidiffusion_sampling.sample_regional_multidiffusion( - model=model, - seed=123, - steps=2, - cfg=7.0, - sampler_name="euler", - scheduler="normal", - positive=[{"model_conds": {}}], - negative=[{"model_conds": {}}], - latent_image=latent_image, - regions=(_region(0, 0, 4, 4, [{"model_conds": {}}], latent_width=8),), - denoise=1.0, - global_prompt_weight=0.25, - ) - - assert output is not latent_image - assert output["samples"] is sampled - assert output["kept"] == "value" - assert "downscale_ratio_spacial" not in output - assert calls["sample_custom"]["model"] is not model - assert calls["sample_custom"]["model"].calc_wrapper is not None - assert calls["sample_custom"]["sampler"] is sampler - assert calls["sample_custom"]["noise"] is fixed_noise - assert calls["sample_custom"]["disable_pbar"] is True - - -def test_sample_accepts_singleton_depth_5d_latent( - monkeypatch: pytest.MonkeyPatch, -) -> None: - """Anima-style singleton-depth latents pass runtime validation.""" - - def fake_sample_custom( - _model: object, - _noise: torch.Tensor, - _cfg: float, - _sampler: object, - _sigmas: torch.Tensor, - _positive: object, - _negative: object, - latent_image: torch.Tensor, - **_kwargs: object, - ) -> torch.Tensor: - """Return a deterministic sampled latent for shape validation.""" - - return latent_image + 1.0 - - model = FakeModel() - sampler = FakeSampler() - latent_samples = torch.zeros((1, 16, 1, 4, 8), dtype=torch.float32) - fixed_sigmas = torch.tensor([1.0, 0.0], dtype=torch.float32) - - 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=None: torch.ones_like(samples), - ) - monkeypatch.setattr(latent_preview, "prepare_callback", lambda _model, _steps: None) - monkeypatch.setattr(comfy_utils, "PROGRESS_BAR_ENABLED", True) - monkeypatch.setattr(comfy_sample, "sample_custom", fake_sample_custom) - - output = regional_multidiffusion_sampling.sample_regional_multidiffusion( - model=model, - seed=123, - steps=2, - cfg=7.0, - sampler_name="euler", - scheduler="normal", - positive=[{"model_conds": {}}], - negative=[{"model_conds": {}}], - latent_image={"samples": latent_samples}, - regions=(_region(0, 0, 4, 4, [{"model_conds": {}}], latent_width=8),), - denoise=1.0, - global_prompt_weight=0.25, - ) - - assert torch.equal(output["samples"], latent_samples + 1.0) - - -def _region( - x: int, - y: int, - width: int, - height: int, - positive: object, - *, - latent_width: int = 4, - latent_height: int = 4, -) -> LatentRegion: - """Return one full-weight latent region.""" - - if isinstance(positive, str): - positive = _raw_conditioning(positive) - mask = torch.zeros((latent_height, latent_width), dtype=torch.float32) - mask[y : y + height, x : x + width] = 1.0 - return LatentRegion( - index=0, - label="region", - latent_box=LatentBox(x, y, width, height), - latent_mask=mask, - positive=positive, - ) - - -def _raw_conditioning(name: str) -> list[list[object]]: - """Return a raw Comfy CONDITIONING-like value with a visible test name.""" - - return [[torch.zeros((1, 1, 1), dtype=torch.float32), {"name": name}]] - - -def _condition_name(conditioning: object) -> str: - """Return the test-visible name from raw, processed, or sentinel conditioning.""" - - if isinstance(conditioning, str): - return conditioning - if isinstance(conditioning, list) and conditioning: - first = conditioning[0] - if isinstance(first, dict): - return str(first.get("name", "")) - if ( - isinstance(first, list | tuple) - and len(first) > 1 - and isinstance( - first[1], - dict, - ) - ): - return str(first[1].get("name", "")) - return "" diff --git a/tests/tooling/__init__.py b/tests/tooling/__init__.py new file mode 100644 index 0000000..b1f60e6 --- /dev/null +++ b/tests/tooling/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own tooling test behavior.""" diff --git a/tests/test_benchmark_anima_regional_lora_cli.py b/tests/tooling/test_benchmark_anima_regional_lora_cli.py similarity index 86% rename from tests/test_benchmark_anima_regional_lora_cli.py rename to tests/tooling/test_benchmark_anima_regional_lora_cli.py index f431bdc..e114695 100644 --- a/tests/test_benchmark_anima_regional_lora_cli.py +++ b/tests/tooling/test_benchmark_anima_regional_lora_cli.py @@ -10,11 +10,13 @@ import subprocess import sys from pathlib import Path +from support.repository import REPOSITORY_ROOT + def test_benchmark_supports_direct_filename_execution(tmp_path: Path) -> None: """Resolve repository and Comfy packages outside either source root.""" - script = Path(__file__).parents[1] / "tools" / "benchmark_anima_regional_lora.py" + script = REPOSITORY_ROOT / "tools" / "benchmark_anima_regional_lora.py" result = subprocess.run( [sys.executable, str(script), "--help"], diff --git a/tests/test_cold_path_capture_probe.py b/tests/tooling/test_cold_path_capture_probe.py similarity index 100% rename from tests/test_cold_path_capture_probe.py rename to tests/tooling/test_cold_path_capture_probe.py diff --git a/tests/test_global_lora_visual_proof.py b/tests/tooling/test_global_lora_visual_proof.py similarity index 100% rename from tests/test_global_lora_visual_proof.py rename to tests/tooling/test_global_lora_visual_proof.py diff --git a/tests/test_indexed_model_call_capture.py b/tests/tooling/test_indexed_model_call_capture.py similarity index 100% rename from tests/test_indexed_model_call_capture.py rename to tests/tooling/test_indexed_model_call_capture.py diff --git a/tests/test_interop_validation_profile.py b/tests/tooling/test_interop_validation_profile.py similarity index 100% rename from tests/test_interop_validation_profile.py rename to tests/tooling/test_interop_validation_profile.py diff --git a/tests/test_materialization_parity_node.py b/tests/tooling/test_materialization_parity_node.py similarity index 100% rename from tests/test_materialization_parity_node.py rename to tests/tooling/test_materialization_parity_node.py diff --git a/tests/test_materialization_parity_probe.py b/tests/tooling/test_materialization_parity_probe.py similarity index 100% rename from tests/test_materialization_parity_probe.py rename to tests/tooling/test_materialization_parity_probe.py diff --git a/tests/test_materialized_variant_comparison.py b/tests/tooling/test_materialized_variant_comparison.py similarity index 100% rename from tests/test_materialized_variant_comparison.py rename to tests/tooling/test_materialized_variant_comparison.py diff --git a/tests/test_preparation_collaborator_profile.py b/tests/tooling/test_preparation_collaborator_profile.py similarity index 99% rename from tests/test_preparation_collaborator_profile.py rename to tests/tooling/test_preparation_collaborator_profile.py index 936f918..eed9019 100644 --- a/tests/test_preparation_collaborator_profile.py +++ b/tests/tooling/test_preparation_collaborator_profile.py @@ -13,6 +13,9 @@ from typing import Any, cast import pytest import torch +from simple_syrup.domain.attention_coupling_preparation import ( + AttentionCouplingPreparation, +) from simple_syrup.domain.processed_regional_attention import ( ProcessedRegionalAttentionPlan, ) @@ -38,7 +41,6 @@ from simple_syrup.services.attention_coupling_model_family_selector import ( AttentionCouplingModelFamilySelector, ) from simple_syrup.services.attention_coupling_preparation_service import ( - AttentionCouplingPreparation, AttentionCouplingPreparationService, ) from tools.attention_coupling_benchmark.comfy_probe import ( diff --git a/tests/test_python_call_profile.py b/tests/tooling/test_python_call_profile.py similarity index 100% rename from tests/test_python_call_profile.py rename to tests/tooling/test_python_call_profile.py diff --git a/tests/test_run_managed_comfy_baseline.py b/tests/tooling/test_run_managed_comfy_baseline.py similarity index 100% rename from tests/test_run_managed_comfy_baseline.py rename to tests/tooling/test_run_managed_comfy_baseline.py diff --git a/tests/test_run_sdxl_attention_couple_parity.py b/tests/tooling/test_run_sdxl_attention_couple_parity.py similarity index 97% rename from tests/test_run_sdxl_attention_couple_parity.py rename to tests/tooling/test_run_sdxl_attention_couple_parity.py index ed76de9..51dd2d7 100644 --- a/tests/test_run_sdxl_attention_couple_parity.py +++ b/tests/tooling/test_run_sdxl_attention_couple_parity.py @@ -13,7 +13,9 @@ from types import SimpleNamespace from PIL import Image from pytest import MonkeyPatch -from sdxl_visual_test_inventory import visual_inventory +from regional_generation.sdxl.support.sdxl_visual_test_inventory import ( + visual_inventory, +) import tools.run_sdxl_attention_couple_parity as runner from tools.comfy_api import ImageReference, JsonObject diff --git a/tests/test_run_sdxl_attention_coupling_integration.py b/tests/tooling/test_run_sdxl_attention_coupling_integration.py similarity index 100% rename from tests/test_run_sdxl_attention_coupling_integration.py rename to tests/tooling/test_run_sdxl_attention_coupling_integration.py diff --git a/tests/test_run_sdxl_global_lora_reference.py b/tests/tooling/test_run_sdxl_global_lora_reference.py similarity index 97% rename from tests/test_run_sdxl_global_lora_reference.py rename to tests/tooling/test_run_sdxl_global_lora_reference.py index 1db707a..d647bbb 100644 --- a/tests/test_run_sdxl_global_lora_reference.py +++ b/tests/tooling/test_run_sdxl_global_lora_reference.py @@ -13,7 +13,9 @@ from types import SimpleNamespace from PIL import Image from pytest import MonkeyPatch -from sdxl_visual_test_inventory import visual_inventory +from regional_generation.sdxl.support.sdxl_visual_test_inventory import ( + visual_inventory, +) import tools.run_sdxl_global_lora_reference as runner from tools.comfy_api import ImageReference, JsonObject diff --git a/tests/test_run_sdxl_post_optimization_visual_proof.py b/tests/tooling/test_run_sdxl_post_optimization_visual_proof.py similarity index 100% rename from tests/test_run_sdxl_post_optimization_visual_proof.py rename to tests/tooling/test_run_sdxl_post_optimization_visual_proof.py diff --git a/tests/test_run_sdxl_regional_lora_visual_matrix.py b/tests/tooling/test_run_sdxl_regional_lora_visual_matrix.py similarity index 98% rename from tests/test_run_sdxl_regional_lora_visual_matrix.py rename to tests/tooling/test_run_sdxl_regional_lora_visual_matrix.py index 64af150..6d132fa 100644 --- a/tests/test_run_sdxl_regional_lora_visual_matrix.py +++ b/tests/tooling/test_run_sdxl_regional_lora_visual_matrix.py @@ -9,7 +9,9 @@ from __future__ import annotations from pathlib import Path import pytest -from sdxl_visual_test_inventory import visual_inventory +from regional_generation.sdxl.support.sdxl_visual_test_inventory import ( + visual_inventory, +) import tools.sdxl_attention_coupling_integration.visual_matrix_execution as execution from tools.comfy_api import JsonObject diff --git a/tests/test_selected_external_visual_runner.py b/tests/tooling/test_selected_external_visual_runner.py similarity index 97% rename from tests/test_selected_external_visual_runner.py rename to tests/tooling/test_selected_external_visual_runner.py index f4b081b..46b9fca 100644 --- a/tests/test_selected_external_visual_runner.py +++ b/tests/tooling/test_selected_external_visual_runner.py @@ -9,7 +9,9 @@ from __future__ import annotations from pathlib import Path import pytest -from sdxl_visual_test_inventory import visual_inventory +from regional_generation.sdxl.support.sdxl_visual_test_inventory import ( + visual_inventory, +) from tools.comfy_integration.artifacts import IntegrationArtifacts from tools.sdxl_attention_coupling_integration import ( diff --git a/tests/test_synchronized_phase_timing.py b/tests/tooling/test_synchronized_phase_timing.py similarity index 100% rename from tests/test_synchronized_phase_timing.py rename to tests/tooling/test_synchronized_phase_timing.py diff --git a/tests/test_tensor_snapshot_probe.py b/tests/tooling/test_tensor_snapshot_probe.py similarity index 100% rename from tests/test_tensor_snapshot_probe.py rename to tests/tooling/test_tensor_snapshot_probe.py diff --git a/tests/test_text_encoder_lora_fixture.py b/tests/tooling/test_text_encoder_lora_fixture.py similarity index 100% rename from tests/test_text_encoder_lora_fixture.py rename to tests/tooling/test_text_encoder_lora_fixture.py diff --git a/tests/test_text_encoder_lora_matrix.py b/tests/tooling/test_text_encoder_lora_matrix.py similarity index 100% rename from tests/test_text_encoder_lora_matrix.py rename to tests/tooling/test_text_encoder_lora_matrix.py diff --git a/tests/test_text_encoder_lora_results.py b/tests/tooling/test_text_encoder_lora_results.py similarity index 100% rename from tests/test_text_encoder_lora_results.py rename to tests/tooling/test_text_encoder_lora_results.py diff --git a/tests/test_text_encoder_lora_workflow.py b/tests/tooling/test_text_encoder_lora_workflow.py similarity index 100% rename from tests/test_text_encoder_lora_workflow.py rename to tests/tooling/test_text_encoder_lora_workflow.py diff --git a/tools/add_license_headers.py b/tools/add_license_headers.py index 2effdd6..890ff97 100644 --- a/tools/add_license_headers.py +++ b/tools/add_license_headers.py @@ -76,6 +76,7 @@ def _tracked_source_files() -> list[Path]: capture_output=True, check=True, text=True, + timeout=30.0, ) except subprocess.CalledProcessError as exc: print(f"Error running git ls-files: {exc}", file=sys.stderr) diff --git a/tools/anima_regression_oracle/execution.py b/tools/anima_regression_oracle/execution.py index a1b203e..5fd20bf 100644 --- a/tools/anima_regression_oracle/execution.py +++ b/tools/anima_regression_oracle/execution.py @@ -55,6 +55,7 @@ class OracleCommandExecutor: stdout=stdout, stderr=stderr, text=True, + timeout=1800.0, ) observation = CommandObservation( command.identity, diff --git a/tools/anima_regression_oracle/manifest.py b/tools/anima_regression_oracle/manifest.py index 6268a00..a0031a4 100644 --- a/tools/anima_regression_oracle/manifest.py +++ b/tools/anima_regression_oracle/manifest.py @@ -95,38 +95,38 @@ def default_manifest(repo_root: Path) -> AnimaRegressionManifest: evidence_root / "global-style-character-proof" / "20260812T224908Z-892f5d8f" ) focused_tests = ( - "tests/test_anima_module_surface.py", - "tests/test_anima_lora_weight_categories.py", - "tests/test_anima_lora_linear.py", - "tests/test_anima_multi_lora_fidelity.py", - "tests/test_anima_multi_lora_composition.py", - "tests/test_anima_full_tile_lora_equivalence.py", - "tests/test_regional_lora_conditioning_adapter.py", - "tests/test_anima_lora_block.py", - "tests/test_anima_lora_block_inactive_schedule.py", - "tests/test_anima_projection_batch.py", - "tests/test_anima_lora_combined_multiplier_cache.py", - "tests/test_anima_regional_lora_device_cache_lifecycle.py", - "tests/test_anima_single_adapter_mutations.py", - "tests/test_anima_regional_lora_target_ownership.py", - "tests/test_anima_composition_phase.py", - "tests/test_anima_cross_attention.py", - "tests/test_anima_self_attention_ownership.py", - "tests/test_anima_self_attention_coherence.py", - "tests/test_anima_regional_self_attention.py", - "tests/test_anima_branch_batch.py", - "tests/test_anima_query_masks.py", - "tests/test_anima_query_activity.py", - "tests/test_anima_regional_diagnostics.py", - "tests/test_anima_regional_permutation_diagnostics.py", - "tests/test_anima_attention_device_cache_lifecycle.py", - "tests/test_anima_attention_coupling_integration.py", - "tests/test_anima_tiled_attention_coupling_integration.py", - "tests/test_anima_contextual_attention_coupling_integration.py", - "tests/test_anima_regional_lora_admission_results.py", - "tests/test_anima_regional_lora_performance_manifest.py", - "tests/test_anima_regional_lora_performance_results.py", - "tests/test_anima_regional_lora_scaling_results.py", + "tests/regional_generation/anima/test_anima_module_surface.py", + "tests/regional_generation/anima/test_anima_lora_weight_categories.py", + "tests/regional_generation/anima/test_anima_lora_linear.py", + "tests/regional_generation/anima/test_anima_multi_lora_fidelity.py", + "tests/regional_generation/anima/test_anima_multi_lora_composition.py", + "tests/regional_generation/anima/test_anima_full_tile_lora_equivalence.py", + "tests/regional_generation/regional/test_regional_lora_conditioning_adapter.py", + "tests/regional_generation/anima/test_anima_lora_block.py", + "tests/regional_generation/anima/test_anima_lora_block_inactive_schedule.py", + "tests/regional_generation/anima/test_anima_projection_batch.py", + "tests/regional_generation/anima/test_anima_lora_combined_multiplier_cache.py", + "tests/regional_generation/anima/test_anima_regional_lora_device_cache_lifecycle.py", + "tests/regional_generation/anima/test_anima_single_adapter_mutations.py", + "tests/regional_generation/anima/test_anima_regional_lora_target_ownership.py", + "tests/regional_generation/anima/test_anima_composition_phase.py", + "tests/regional_generation/anima/test_anima_cross_attention.py", + "tests/regional_generation/anima/test_anima_self_attention_ownership.py", + "tests/regional_generation/anima/test_anima_self_attention_coherence.py", + "tests/regional_generation/anima/test_anima_regional_self_attention.py", + "tests/regional_generation/anima/test_anima_branch_batch.py", + "tests/regional_generation/anima/test_anima_query_masks.py", + "tests/regional_generation/anima/test_anima_query_activity.py", + "tests/regional_generation/anima/test_anima_regional_diagnostics.py", + "tests/regional_generation/anima/test_anima_regional_permutation_diagnostics.py", + "tests/regional_generation/anima/test_anima_attention_device_cache_lifecycle.py", + "tests/regional_generation/anima/test_anima_attention_coupling_integration.py", + "tests/regional_generation/anima/test_anima_tiled_attention_coupling_integration.py", + "tests/regional_generation/anima/test_anima_contextual_attention_coupling_integration.py", + "tests/regional_generation/anima/test_anima_regional_lora_admission_results.py", + "tests/regional_generation/anima/test_anima_regional_lora_performance_manifest.py", + "tests/regional_generation/anima/test_anima_regional_lora_performance_results.py", + "tests/regional_generation/anima/test_anima_regional_lora_scaling_results.py", ) return AnimaRegressionManifest( images=( diff --git a/tools/architecture_governance/__init__.py b/tools/architecture_governance/__init__.py new file mode 100644 index 0000000..0a61882 --- /dev/null +++ b/tools/architecture_governance/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Enforce current architecture limits, debt, and waivers.""" diff --git a/tools/architecture_governance/import_boundaries.py b/tools/architecture_governance/import_boundaries.py new file mode 100644 index 0000000..572f326 --- /dev/null +++ b/tools/architecture_governance/import_boundaries.py @@ -0,0 +1,165 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Enforce SimpleSyrup package dependency direction with exact debt snapshots.""" + +from __future__ import annotations + +import ast +import tomllib +from datetime import UTC, date, datetime +from pathlib import Path +from typing import TypedDict, cast + +from .metrics import source_fingerprint +from .model import Diagnostic + +IMPORT_DEBT_PATH = Path("governance/architecture/import_debt.toml") +_FORBIDDEN_EDGES = frozenset( + { + ("domain", "masking"), + ("domain", "integration"), + ("domain", "runtime"), + ("domain", "services"), + ("image", "integration"), + ("masking", "integration"), + ("nodes", "integration"), + ("nodes_v3", "integration"), + ("runtime", "integration"), + ("runtime", "nodes"), + ("runtime", "nodes_v3"), + ("runtime", "services"), + ("services", "integration"), + ("services", "nodes"), + ("services", "nodes_v3"), + ("shared", "domain"), + ("shared", "integration"), + ("shared", "runtime"), + ("shared", "services"), + } +) + + +class _ImportDebt(TypedDict): + """Describe one exact source file's known forbidden imports.""" + + id: str + path: str + imports: list[str] + fingerprint: str + issue: str + review_by: date + problem: str + remediation: str + + +def validate_import_boundaries( + root: Path, *, today: date | None = None +) -> list[Diagnostic]: + """Return diagnostics for unreviewed or stale package dependency debt.""" + + registry_path = root / IMPORT_DEBT_PATH + try: + payload = tomllib.loads(registry_path.read_text(encoding="utf-8")) + records = cast(list[_ImportDebt], payload.get("debts", [])) + except (OSError, tomllib.TOMLDecodeError) as error: + return [Diagnostic("IMPORTSTATE001", IMPORT_DEBT_PATH.as_posix(), str(error))] + current = _forbidden_imports(root) + diagnostics: list[Diagnostic] = [] + recorded_paths: set[str] = set() + current_date = today or datetime.now(UTC).date() + for record in records: + path = record["path"] + if path in recorded_paths: + diagnostics.append( + Diagnostic( + "IMPORTSTATE002", + IMPORT_DEBT_PATH.as_posix(), + f"forbidden-import path {path} has multiple debt records", + ) + ) + continue + recorded_paths.add(path) + expected = tuple(record["imports"]) + if expected != tuple(sorted(expected)) or current.get(path) != expected: + diagnostics.append( + Diagnostic( + "IMPORTDEBT001", + IMPORT_DEBT_PATH.as_posix(), + f"debt {record['id']} must match the exact sorted forbidden " + f"imports for {path}", + ) + ) + elif source_fingerprint(root, (path,)) != record["fingerprint"]: + diagnostics.append( + Diagnostic( + "IMPORTDEBT002", + IMPORT_DEBT_PATH.as_posix(), + f"debt {record['id']} source changed; repeat ownership review", + ) + ) + if record["review_by"] < current_date: + diagnostics.append( + Diagnostic( + "IMPORTDEBT003", + IMPORT_DEBT_PATH.as_posix(), + f"debt {record['id']} expired on {record['review_by'].isoformat()}", + ) + ) + for path, imports in current.items(): + if path not in recorded_paths: + diagnostics.append( + Diagnostic( + "IMPORT001", + path, + "forbidden package dependency requires an exact debt record: " + + ", ".join(imports), + ) + ) + for path in recorded_paths - current.keys(): + diagnostics.append( + Diagnostic( + "IMPORTDEBT004", + IMPORT_DEBT_PATH.as_posix(), + f"recorded forbidden-import debt for {path} is resolved or stale", + ) + ) + return sorted(diagnostics, key=lambda item: (item.path, item.rule, item.message)) + + +def _forbidden_imports(root: Path) -> dict[str, tuple[str, ...]]: + """Discover exact first-party imports that point against layer direction.""" + + package_root = root / "simple_syrup" + findings: dict[str, tuple[str, ...]] = {} + for path in sorted(package_root.rglob("*.py")): + if "third_party" in path.parts: + continue + relative = path.relative_to(package_root) + source_layer = relative.parts[0] + modules = { + module + for node in ast.walk(ast.parse(path.read_text(encoding="utf-8"))) + for module in _imported_modules(node, relative) + if module.startswith("simple_syrup.") + and (source_layer, module.split(".", maxsplit=2)[1]) in _FORBIDDEN_EDGES + } + if modules: + findings[path.relative_to(root).as_posix()] = tuple(sorted(modules)) + return findings + + +def _imported_modules(node: ast.AST, relative_path: Path) -> tuple[str, ...]: + """Resolve absolute and package-relative imports for one syntax node.""" + + if isinstance(node, ast.Import): + return tuple(alias.name for alias in node.names) + if not isinstance(node, ast.ImportFrom): + return () + if node.level == 0: + return (node.module,) if node.module else () + package_parts = list(relative_path.with_suffix("").parts[:-1]) + retained = package_parts[: len(package_parts) - (node.level - 1)] + suffix = node.module.split(".") if node.module else [] + return (".".join(("simple_syrup", *retained, *suffix)),) diff --git a/tools/architecture_governance/loading.py b/tools/architecture_governance/loading.py new file mode 100644 index 0000000..384e72f --- /dev/null +++ b/tools/architecture_governance/loading.py @@ -0,0 +1,246 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Load strict machine-readable architecture policy and current state.""" + +from __future__ import annotations + +import tomllib +from collections.abc import Mapping +from datetime import date +from pathlib import Path +from typing import cast + +from .model import ( + ArchitectureDebt, + ArchitecturePolicy, + ArchitectureState, + ArchitectureWaiver, +) + + +def load_policy(path: Path) -> ArchitecturePolicy: + """Load one schema-versioned architecture policy.""" + + data = _document(path, expected_version=2) + _exact_keys(data, {"schema_version", "structure", "registries"}, str(path)) + structure = _mapping(data, "structure") + registries = _mapping(data, "registries") + _exact_keys( + structure, + { + "soft_lines", + "hard_lines", + "source_roots", + "source_files", + "source_extensions", + "excluded_paths", + }, + "structure policy", + ) + _exact_keys( + registries, + {"debt", "waivers"}, + "architecture registries", + ) + return ArchitecturePolicy( + soft_lines=_positive_int(structure, "soft_lines"), + hard_lines=_positive_int(structure, "hard_lines"), + source_roots=tuple( + Path(value) for value in _strings(structure, "source_roots") + ), + source_files=tuple( + Path(value) for value in _string_list(structure, "source_files") + ), + source_extensions=frozenset(_strings(structure, "source_extensions")), + excluded_paths=frozenset(_string_list(structure, "excluded_paths")), + debt_registry=Path(_string(registries, "debt")), + waiver_registry=Path(_string(registries, "waivers")), + ) + + +def load_state(root: Path, policy: ArchitecturePolicy) -> ArchitectureState: + """Load all exact current architecture state registries.""" + + debt_data = _registry_document(root / policy.debt_registry, "debts") + waiver_data = _registry_document(root / policy.waiver_registry, "waivers") + return ArchitectureState( + debts=tuple(_parse_debt(item) for item in _tables(debt_data, "debts")), + waivers=tuple(_parse_waiver(item) for item in _tables(waiver_data, "waivers")), + ) + + +def _parse_debt(data: Mapping[str, object]) -> ArchitectureDebt: + """Parse one assessed mixed-ownership record.""" + + _exact_keys( + data, + { + "id", + "owner", + "paths", + "fingerprint", + "issue", + "review_by", + "responsibilities", + "next_extraction", + }, + "architecture debt", + ) + responsibilities = _strings(data, "responsibilities") + if len(responsibilities) < 2: + raise ValueError("architecture debt must name at least two responsibilities") + return ArchitectureDebt( + identifier=_string(data, "id"), + owner=_string(data, "owner"), + paths=_strings(data, "paths"), + fingerprint=_string(data, "fingerprint"), + issue=_string(data, "issue"), + review_by=_date(data, "review_by"), + responsibilities=responsibilities, + next_extraction=_string(data, "next_extraction"), + ) + + +def _parse_waiver(data: Mapping[str, object]) -> ArchitectureWaiver: + """Parse one bounded structural or remediation waiver.""" + + kind = _string(data, "kind") + if kind not in {"structural", "remediation"}: + raise ValueError("architecture waiver kind must be structural or remediation") + fields = { + "id", + "owner", + "rule", + "path", + "kind", + "justification", + "issue", + "review_by", + "max_lines", + } + if kind == "remediation": + fields |= {"next_limit", "debt"} + _exact_keys(data, fields, "architecture waiver") + return ArchitectureWaiver( + identifier=_string(data, "id"), + owner=_string(data, "owner"), + rule=_string(data, "rule"), + path=_string(data, "path"), + kind=kind, + justification=_string(data, "justification"), + issue=_string(data, "issue"), + review_by=_date(data, "review_by"), + max_lines=_positive_int(data, "max_lines"), + next_limit=( + _positive_int(data, "next_limit") if kind == "remediation" else None + ), + debt=_string(data, "debt") if kind == "remediation" else None, + ) + + +def _document(path: Path, *, expected_version: int) -> Mapping[str, object]: + """Read one TOML document and validate its schema envelope.""" + + data = cast(Mapping[str, object], tomllib.loads(path.read_text(encoding="utf-8"))) + if data.get("schema_version") != expected_version: + raise ValueError(f"{path} must declare schema_version = {expected_version}") + return data + + +def _registry_document(path: Path, collection: str) -> Mapping[str, object]: + """Read one strict registry envelope containing current state only.""" + + data = _document(path, expected_version=1) + _exact_keys(data, {"schema_version", collection}, str(path)) + return data + + +def _exact_keys( + data: Mapping[str, object], + expected: set[str], + label: str, +) -> None: + """Reject missing fields and history-shaped surplus fields.""" + + missing = expected - set(data) + unknown = set(data) - expected + if missing or unknown: + raise ValueError( + f"{label} fields differ: missing={sorted(missing)}, " + f"unsupported={sorted(unknown)}" + ) + + +def _mapping(data: Mapping[str, object], key: str) -> Mapping[str, object]: + """Return one required TOML table.""" + + value = data.get(key) + if not isinstance(value, dict): + raise TypeError(f"{key} must be a table") + return cast(Mapping[str, object], value) + + +def _tables(data: Mapping[str, object], key: str) -> tuple[Mapping[str, object], ...]: + """Return one optional array of TOML tables.""" + + value = data.get(key, []) + if not isinstance(value, list) or not all(isinstance(item, dict) for item in value): + raise TypeError(f"{key} must be an array of tables") + return tuple(cast(list[Mapping[str, object]], value)) + + +def _string(data: Mapping[str, object], key: str) -> str: + """Return one required nonempty string.""" + + value = data.get(key) + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{key} must be a nonempty string") + return value + + +def _strings(data: Mapping[str, object], key: str) -> tuple[str, ...]: + """Return one required nonempty array of unique strings.""" + + value = data.get(key) + if not isinstance(value, list) or not value: + raise ValueError(f"{key} must be a nonempty string array") + if not all(isinstance(item, str) and item.strip() for item in value): + raise ValueError(f"{key} must contain nonempty strings") + strings = tuple(cast(list[str], value)) + if len(strings) != len(set(strings)): + raise ValueError(f"{key} must not contain duplicates") + return strings + + +def _string_list(data: Mapping[str, object], key: str) -> tuple[str, ...]: + """Return one required array of unique nonempty strings.""" + + value = data.get(key) + if not isinstance(value, list): + raise TypeError(f"{key} must be a string array") + if not all(isinstance(item, str) and item.strip() for item in value): + raise ValueError(f"{key} must contain nonempty strings") + strings = tuple(cast(list[str], value)) + if len(strings) != len(set(strings)): + raise ValueError(f"{key} must not contain duplicates") + return strings + + +def _positive_int(data: Mapping[str, object], key: str) -> int: + """Return one required positive integer.""" + + value = data.get(key) + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"{key} must be a positive integer") + return value + + +def _date(data: Mapping[str, object], key: str) -> date: + """Return one required TOML date.""" + + value = data.get(key) + if not isinstance(value, date): + raise TypeError(f"{key} must be an ISO date") + return value diff --git a/tools/architecture_governance/metrics.py b/tools/architecture_governance/metrics.py new file mode 100644 index 0000000..8dbdc2b --- /dev/null +++ b/tools/architecture_governance/metrics.py @@ -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 + +"""Own stable source measurements used by architecture governance.""" + +from __future__ import annotations + +import hashlib +from collections.abc import Iterable +from pathlib import Path + +from .model import ArchitecturePolicy + + +def production_line_count(path: Path) -> int: + """Count nonblank physical lines that are not language comment-only lines.""" + + lines = normalized_source(path).splitlines() + if path.suffix in {".js", ".mjs", ".cjs", ".ts"}: + return _javascript_production_line_count(lines) + return sum( + 1 for line in lines if line.strip() and not line.lstrip().startswith("#") + ) + + +def source_fingerprint(root: Path, paths: Iterable[str]) -> str: + """Return a newline-stable fingerprint for exact repository paths.""" + + digest = hashlib.sha256() + for relative_path in sorted(paths): + normalized_path = relative_path.replace("\\", "/") + digest.update(normalized_path.encode("utf-8")) + digest.update(b"\0") + digest.update(normalized_source(root / normalized_path).encode("utf-8")) + digest.update(b"\0") + return f"sha256:{digest.hexdigest()}" + + +def normalized_source(path: Path) -> str: + """Read UTF-8 source with canonical newlines for cross-platform identity.""" + + return path.read_text(encoding="utf-8").replace("\r\n", "\n").replace("\r", "\n") + + +def governed_source_paths( + root: Path, + policy: ArchitecturePolicy, +) -> tuple[Path, ...]: + """Return exact authored-code paths declared by the architecture policy.""" + + rooted_candidates = { + path + for source_root in policy.source_roots + for path in (root / source_root).rglob("*") + if path.is_file() and path.suffix in policy.source_extensions + } + exact_candidates = { + path + for source_file in policy.source_files + if (path := root / source_file).is_file() + and path.suffix in policy.source_extensions + } + return tuple( + path + for path in sorted(rooted_candidates | exact_candidates) + if path.relative_to(root).as_posix() not in policy.excluded_paths + and "__pycache__" not in path.parts + ) + + +def _javascript_production_line_count(lines: list[str]) -> int: + """Count JavaScript source while excluding standalone line and block comments.""" + + count = 0 + inside_block_comment = False + for line in lines: + remaining = line.strip() + if not remaining: + continue + while remaining: + if inside_block_comment: + if "*/" not in remaining: + break + inside_block_comment = False + remaining = remaining.split("*/", maxsplit=1)[1].strip() + continue + if remaining.startswith("//"): + break + if remaining.startswith("/*"): + if "*/" not in remaining[2:]: + inside_block_comment = True + break + remaining = remaining.split("*/", maxsplit=1)[1].strip() + continue + count += 1 + break + return count diff --git a/tools/architecture_governance/model.py b/tools/architecture_governance/model.py new file mode 100644 index 0000000..5af3d7d --- /dev/null +++ b/tools/architecture_governance/model.py @@ -0,0 +1,79 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define immutable architecture governance values.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import date +from pathlib import Path + + +@dataclass(frozen=True, slots=True) +class Diagnostic: + """Describe one actionable architecture policy result.""" + + rule: str + path: str + message: str + severity: str = "error" + + def render(self) -> str: + """Render a stable path-oriented diagnostic.""" + + return f"{self.path}:1: {self.severity} {self.rule}: {self.message}" + + +@dataclass(frozen=True, slots=True) +class ArchitecturePolicy: + """Define the source scope and structural limits.""" + + soft_lines: int + hard_lines: int + source_roots: tuple[Path, ...] + source_files: tuple[Path, ...] + source_extensions: frozenset[str] + excluded_paths: frozenset[str] + debt_registry: Path + waiver_registry: Path + + +@dataclass(frozen=True, slots=True) +class ArchitectureDebt: + """Describe current assessed mixed ownership and its next extraction.""" + + identifier: str + owner: str + paths: tuple[str, ...] + fingerprint: str + issue: str + review_by: date + responsibilities: tuple[str, ...] + next_extraction: str + + +@dataclass(frozen=True, slots=True) +class ArchitectureWaiver: + """Describe one exact, bounded structural gate exception.""" + + identifier: str + owner: str + rule: str + path: str + kind: str + justification: str + issue: str + review_by: date + max_lines: int + next_limit: int | None + debt: str | None + + +@dataclass(frozen=True, slots=True) +class ArchitectureState: + """Collect current debt and waiver snapshots.""" + + debts: tuple[ArchitectureDebt, ...] + waivers: tuple[ArchitectureWaiver, ...] diff --git a/tools/architecture_governance/soft_reviews.py b/tools/architecture_governance/soft_reviews.py new file mode 100644 index 0000000..7126211 --- /dev/null +++ b/tools/architecture_governance/soft_reviews.py @@ -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 + +"""Validate exact human review of every current soft-ceiling source file.""" + +from __future__ import annotations + +import tomllib +from datetime import UTC, date, datetime +from pathlib import Path +from typing import TypedDict, cast + +from .loading import load_policy +from .metrics import governed_source_paths, production_line_count, source_fingerprint +from .model import Diagnostic + +SOFT_REVIEW_PATH = Path("governance/architecture/soft_reviews.toml") + + +class _Remediation(TypedDict): + """Describe one reviewed soft-ceiling mixed-ownership finding.""" + + id: str + path: str + responsibilities: list[str] + next_extraction: str + + +def validate_soft_reviews(root: Path, *, today: date | None = None) -> list[Diagnostic]: + """Require an exact current disposition for every soft-ceiling file.""" + + registry = root / SOFT_REVIEW_PATH + try: + payload = tomllib.loads(registry.read_text(encoding="utf-8")) + policy = load_policy(root / "governance/architecture/policy.toml") + cohesive = tuple(cast(list[str], payload["cohesive_paths"])) + debt = tuple(cast(list[str], payload["debt_paths"])) + remediations = cast(list[_Remediation], payload["remediations"]) + fingerprint = cast(str, payload["fingerprint"]) + review_by = cast(date, payload["review_by"]) + except (KeyError, OSError, TypeError, ValueError, tomllib.TOMLDecodeError) as error: + return [Diagnostic("SOFTSTATE001", SOFT_REVIEW_PATH.as_posix(), str(error))] + current = tuple( + path.relative_to(root).as_posix() + for path in governed_source_paths(root, policy) + if policy.soft_lines < production_line_count(path) <= policy.hard_lines + ) + diagnostics: list[Diagnostic] = [] + reviewed = tuple(sorted((*cohesive, *debt))) + if cohesive != tuple(sorted(cohesive)) or debt != tuple(sorted(debt)): + diagnostics.append( + Diagnostic( + "SOFTSTATE002", + SOFT_REVIEW_PATH.as_posix(), + "cohesive_paths and debt_paths must use stable sorted order", + ) + ) + if set(cohesive) & set(debt) or reviewed != current: + diagnostics.append( + Diagnostic( + "SOFTSTATE003", + SOFT_REVIEW_PATH.as_posix(), + "soft reviews must classify every exact current warning once", + ) + ) + elif source_fingerprint(root, reviewed) != fingerprint: + diagnostics.append( + Diagnostic( + "SOFTSTATE004", + SOFT_REVIEW_PATH.as_posix(), + "soft-ceiling source changed; repeat human ownership review", + ) + ) + if review_by < (today or datetime.now(UTC).date()): + diagnostics.append( + Diagnostic( + "SOFTSTATE005", + SOFT_REVIEW_PATH.as_posix(), + f"soft-ceiling review expired on {review_by.isoformat()}", + ) + ) + remediation_paths = tuple(sorted(item["path"] for item in remediations)) + if remediation_paths != debt or len(remediation_paths) != len( + set(remediation_paths) + ): + diagnostics.append( + Diagnostic( + "SOFTDEBT001", + SOFT_REVIEW_PATH.as_posix(), + "every soft debt path requires exactly one remediation record", + ) + ) + for item in remediations: + if len(item["responsibilities"]) < 2 or len(item["next_extraction"]) < 80: + diagnostics.append( + Diagnostic( + "SOFTDEBT002", + SOFT_REVIEW_PATH.as_posix(), + f"remediation {item['id']} requires substantive source review", + ) + ) + return diagnostics diff --git a/tools/architecture_governance/validation.py b/tools/architecture_governance/validation.py new file mode 100644 index 0000000..bf449e8 --- /dev/null +++ b/tools/architecture_governance/validation.py @@ -0,0 +1,382 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Validate structural limits and current architecture state.""" + +from __future__ import annotations + +from collections import Counter +from datetime import UTC, date, datetime +from pathlib import Path + +from .loading import load_policy, load_state +from .metrics import governed_source_paths, production_line_count, source_fingerprint +from .model import ( + ArchitectureDebt, + ArchitecturePolicy, + ArchitectureState, + ArchitectureWaiver, + Diagnostic, +) + +_OWNERSHIP_RESPONSE = ( + "Assess the file's concern, authoritative state owner, dependency direction, " + "public boundary, behavior contract, and change cadence. If ownership is " + "mixed, adding behavior is prohibited: characterize the touched behavior, " + "extract focused owners, migrate every caller, and remove replaced code and " + "bridges. A waiver can bound a size exception; it cannot authorize new " + "behavior in mixed code." +) +_MINIMUM_STRUCTURAL_JUSTIFICATION_LENGTH = 160 +_MECHANICAL_STRUCTURAL_JUSTIFICATION = ( + "Its production types and helpers serve that one authoritative contract" +) + + +def validate_repository( + root: Path, + *, + policy_path: Path | None = None, + today: date | None = None, +) -> list[Diagnostic]: + """Return every architecture diagnostic for the current repository.""" + + try: + policy = load_policy( + policy_path or root / "governance/architecture/policy.toml" + ) + state = load_state(root, policy) + except (OSError, TypeError, ValueError) as error: + return [ + Diagnostic("STATE001", "governance/architecture/policy.toml", str(error)) + ] + current_date = today or datetime.now(UTC).date() + diagnostics = [ + *_validate_policy(root, policy), + *_validate_state(root, policy, state, current_date), + ] + diagnostics.extend(_validate_structure(root, policy, state, current_date)) + return sorted( + diagnostics, + key=lambda item: (item.path, item.rule, item.severity, item.message), + ) + + +def _validate_policy(root: Path, policy: ArchitecturePolicy) -> list[Diagnostic]: + """Validate coherent limits and exact configured source scope.""" + + diagnostics: list[Diagnostic] = [] + if policy.soft_lines >= policy.hard_lines: + diagnostics.append( + Diagnostic( + "POLICY001", + "governance/architecture/policy.toml", + "soft_lines must be lower than hard_lines", + ) + ) + for source_root in policy.source_roots: + if not (root / source_root).is_dir(): + diagnostics.append( + Diagnostic( + "POLICY002", + "governance/architecture/policy.toml", + f"source root {source_root.as_posix()} does not exist", + ) + ) + for source_file in policy.source_files: + if not (root / source_file).is_file(): + diagnostics.append( + Diagnostic( + "POLICY004", + "governance/architecture/policy.toml", + f"source file {source_file.as_posix()} does not exist", + ) + ) + elif source_file.suffix not in policy.source_extensions: + diagnostics.append( + Diagnostic( + "POLICY005", + "governance/architecture/policy.toml", + f"source file {source_file.as_posix()} has an ungoverned extension", + ) + ) + for excluded_path in sorted(policy.excluded_paths): + if not (root / excluded_path).is_file(): + diagnostics.append( + Diagnostic( + "POLICY003", + "governance/architecture/policy.toml", + f"excluded source {excluded_path} does not exist", + ) + ) + return diagnostics + + +def _validate_state( + root: Path, + policy: ArchitecturePolicy, + state: ArchitectureState, + today: date, +) -> list[Diagnostic]: + """Validate registry uniqueness, paths, fingerprints, dates, and links.""" + + diagnostics = _validate_unique_state(state) + governed = { + path.relative_to(root).as_posix() + for path in governed_source_paths(root, policy) + } + debt_by_id = {debt.identifier: debt for debt in state.debts} + for debt in state.debts: + registry = policy.debt_registry.as_posix() + valid_paths = tuple(path for path in debt.paths if path in governed) + if len(valid_paths) != len(debt.paths): + diagnostics.append( + Diagnostic( + "DEBT001", + registry, + f"debt {debt.identifier} must reference exact governed " + "source paths", + ) + ) + elif source_fingerprint(root, debt.paths) != debt.fingerprint: + diagnostics.append( + Diagnostic( + "DEBT002", + registry, + f"debt {debt.identifier} no longer matches assessed source; " + "reassess current responsibilities and the next extraction or " + "delete resolved debt", + ) + ) + if debt.review_by < today: + diagnostics.append( + Diagnostic( + "DEBT003", + registry, + f"debt {debt.identifier} review deadline expired on " + f"{debt.review_by.isoformat()}", + ) + ) + linked_waivers = tuple( + waiver + for waiver in state.waivers + if waiver.kind == "remediation" and waiver.debt == debt.identifier + ) + if len(linked_waivers) != 1: + diagnostics.append( + Diagnostic( + "DEBT004", + registry, + f"debt {debt.identifier} must have exactly one linked remediation " + f"waiver; found {len(linked_waivers)}", + ) + ) + for waiver in state.waivers: + diagnostics.extend( + _validate_waiver(root, policy, waiver, debt_by_id, governed, today) + ) + return diagnostics + + +def _validate_unique_state(state: ArchitectureState) -> list[Diagnostic]: + """Reject duplicate paths and record identifiers across current state.""" + + diagnostics: list[Diagnostic] = [] + identifiers = [ + *(debt.identifier for debt in state.debts), + *(waiver.identifier for waiver in state.waivers), + ] + for identifier, count in Counter(identifiers).items(): + if count > 1: + diagnostics.append( + Diagnostic( + "STATE002", + "governance/architecture/policy.toml", + f"architecture record id {identifier} is not unique", + ) + ) + classified_paths = [waiver.path for waiver in state.waivers] + for path, count in Counter(classified_paths).items(): + if count > 1: + diagnostics.append( + Diagnostic( + "STATE003", + "governance/architecture/policy.toml", + f"source path {path} has multiple structural dispositions", + ) + ) + debt_paths = [path for debt in state.debts for path in debt.paths] + for path, count in Counter(debt_paths).items(): + if count > 1: + diagnostics.append( + Diagnostic( + "STATE004", + "governance/architecture/debt.toml", + f"source path {path} appears in multiple debt records", + ) + ) + structural_justifications = [ + waiver.justification for waiver in state.waivers if waiver.kind == "structural" + ] + for _justification, count in Counter(structural_justifications).items(): + if count > 1: + diagnostics.append( + Diagnostic( + "WAIVER011", + "governance/architecture/waivers.toml", + "structural waiver rationales must be unique and file-specific", + ) + ) + return diagnostics + + +def _validate_waiver( + root: Path, + policy: ArchitecturePolicy, + waiver: ArchitectureWaiver, + debts: dict[str, ArchitectureDebt], + governed: set[str], + today: date, +) -> list[Diagnostic]: + """Validate one exact, bounded, current architecture waiver.""" + + registry = policy.waiver_registry.as_posix() + diagnostics: list[Diagnostic] = [] + if waiver.rule != "STRUCT003": + diagnostics.append( + Diagnostic( + "WAIVER001", + registry, + f"waiver {waiver.identifier} uses an unsupported rule", + ) + ) + if waiver.path not in governed: + diagnostics.append( + Diagnostic( + "WAIVER002", + registry, + f"waiver {waiver.identifier} path is not governed", + ) + ) + return diagnostics + lines = production_line_count(root / waiver.path) + if waiver.review_by < today: + diagnostics.append( + Diagnostic( + "WAIVER003", + registry, + f"waiver {waiver.identifier} expired on {waiver.review_by.isoformat()}", + ) + ) + if lines > waiver.max_lines: + diagnostics.append( + Diagnostic( + "WAIVER004", + registry, + f"waiver {waiver.identifier} caps {waiver.path} at " + f"{waiver.max_lines} lines but current source has {lines}", + ) + ) + if lines <= policy.hard_lines: + diagnostics.append( + Diagnostic( + "WAIVER005", + registry, + f"waiver {waiver.identifier} matches no current hard-size " + "finding; delete it", + ) + ) + if lines != waiver.max_lines: + diagnostics.append( + Diagnostic( + "WAIVER008", + registry, + f"waiver {waiver.identifier} must cap the exact current size " + f"of {waiver.path} at {lines} lines; reassess the source", + ) + ) + if waiver.kind == "remediation": + debt = debts.get(waiver.debt or "") + if debt is None or waiver.path not in debt.paths: + diagnostics.append( + Diagnostic( + "WAIVER006", + registry, + f"waiver {waiver.identifier} must link assessed debt for the " + "same exact path", + ) + ) + if waiver.next_limit is None or waiver.next_limit >= waiver.max_lines: + diagnostics.append( + Diagnostic( + "WAIVER007", + registry, + f"waiver {waiver.identifier} next_limit must be lower than " + "max_lines", + ) + ) + else: + if any(waiver.path in debt.paths for debt in debts.values()): + diagnostics.append( + Diagnostic( + "WAIVER009", + registry, + f"structural waiver {waiver.identifier} cannot cover a path " + "recorded as mixed-responsibility debt", + ) + ) + if ( + len(waiver.justification) < _MINIMUM_STRUCTURAL_JUSTIFICATION_LENGTH + or _MECHANICAL_STRUCTURAL_JUSTIFICATION in waiver.justification + ): + diagnostics.append( + Diagnostic( + "WAIVER010", + registry, + f"structural waiver {waiver.identifier} requires a substantive, " + "file-specific cohesion rationale; human source review is " + "mandatory", + ) + ) + return diagnostics + + +def _validate_structure( + root: Path, + policy: ArchitecturePolicy, + state: ArchitectureState, + today: date, +) -> list[Diagnostic]: + """Enforce the structural size ceiling against exact current dispositions.""" + + waivers = { + waiver.path: waiver + for waiver in state.waivers + if waiver.review_by >= today and waiver.rule == "STRUCT003" + } + diagnostics: list[Diagnostic] = [] + for path in governed_source_paths(root, policy): + relative_path = path.relative_to(root).as_posix() + lines = production_line_count(path) + if lines > policy.hard_lines: + if relative_path not in waivers: + diagnostics.append( + Diagnostic( + "STRUCT003", + relative_path, + f"{lines} production lines exceed the hard gate " + f"{policy.hard_lines}. {_OWNERSHIP_RESPONSE}", + ) + ) + elif lines > policy.soft_lines: + diagnostics.append( + Diagnostic( + "STRUCT002", + relative_path, + f"{lines} production lines exceed the soft ceiling " + f"{policy.soft_lines}; assess ownership before extending this file", + severity="warning", + ) + ) + return diagnostics diff --git a/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py b/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py index f8d73d9..a9d9ec4 100644 --- a/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py +++ b/tools/attention_coupling_benchmark/comfy_probe/preparation_collaborator_profile.py @@ -10,6 +10,9 @@ from typing import ClassVar import torch +from simple_syrup.domain.attention_coupling_preparation import ( + AttentionCouplingPreparation, +) from simple_syrup.domain.processed_regional_attention import ( ProcessedRegionalAttentionPlan, ) @@ -38,7 +41,6 @@ from simple_syrup.services.attention_coupling_model_preparation_service import ( AttentionCouplingModelPreparationService, ) from simple_syrup.services.attention_coupling_preparation_service import ( - AttentionCouplingPreparation, AttentionCouplingPreparationService, ) diff --git a/tools/check_architecture.py b/tools/check_architecture.py new file mode 100644 index 0000000..13958f2 --- /dev/null +++ b/tools/check_architecture.py @@ -0,0 +1,43 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Check repository architecture governance.""" + +from __future__ import annotations + +import sys +from collections.abc import Sequence +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from tools.architecture_governance.import_boundaries import validate_import_boundaries +from tools.architecture_governance.soft_reviews import validate_soft_reviews +from tools.architecture_governance.validation import validate_repository + + +def main(argv: Sequence[str] | None = None) -> int: + """Validate current architecture state.""" + + if argv: + raise ValueError("The architecture checker accepts no arguments") + root = Path(__file__).resolve().parents[1] + diagnostics = [ + *validate_repository(root), + *validate_import_boundaries(root), + *validate_soft_reviews(root), + ] + for diagnostic in diagnostics: + print(diagnostic.render()) + errors = [item for item in diagnostics if item.severity == "error"] + if errors: + print(f"FAILED: Found {len(errors)} architecture governance errors.") + return 1 + warning_count = len(diagnostics) + print(f"SUCCESS: Architecture is valid ({warning_count} structural warnings).") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/check_test_governance.py b/tools/check_test_governance.py new file mode 100644 index 0000000..e470fe0 --- /dev/null +++ b/tools/check_test_governance.py @@ -0,0 +1,39 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Check deterministic test governance and reviewed execution state.""" + +from __future__ import annotations + +import sys +from collections.abc import Sequence +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from tools.test_governance.validation import validate_test_governance + + +def main(argv: Sequence[str] | None = None) -> int: + """Validate the repository's exact current test-governance state.""" + + if argv: + raise ValueError("The test-governance checker accepts no arguments") + root = Path(__file__).resolve().parents[1] + result = validate_test_governance(root) + for diagnostic in result.diagnostics: + print(diagnostic.render()) + errors = [item for item in result.diagnostics if item.severity == "error"] + if errors: + print(f"FAILED: Found {len(errors)} test-governance errors.") + return 1 + print( + "SUCCESS: Test governance is valid " + f"({len(result.candidates)} reviewed candidates)." + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/comfy_integration/loopback_port.py b/tools/comfy_integration/loopback_port.py index 94f603a..5b76f2d 100644 --- a/tools/comfy_integration/loopback_port.py +++ b/tools/comfy_integration/loopback_port.py @@ -2,7 +2,7 @@ # Copyright (C) 2026 Artificial Sweetener and contributors # SPDX-License-Identifier: AGPL-3.0-or-later -"""Select and validate unused non-default loopback TCP ports.""" +"""Reserve and validate non-default loopback TCP ports.""" from __future__ import annotations @@ -11,16 +11,50 @@ import socket PROTECTED_COMFY_PORTS = frozenset({8188, 8297}) -def select_unused_loopback_port() -> int: - """Return one unused ephemeral port outside protected Comfy instances.""" +class LoopbackPortReservation: + """Own an exclusive loopback socket until the managed launcher is ready.""" + + def __init__(self, listener: socket.socket) -> None: + """Retain the bound socket and its OS-assigned port.""" + + self._listener: socket.socket | None = listener + self.port = int(listener.getsockname()[1]) + + def release(self) -> None: + """Release the reservation immediately before the child launch attempt.""" + + listener = self._listener + if listener is not None: + self._listener = None + listener.close() + + def __enter__(self) -> LoopbackPortReservation: + """Expose the live reservation to its lifecycle owner.""" + + return self + + def __exit__(self, exc_type: object, exc: object, traceback: object) -> None: + """Release the reservation on every owner exit path.""" + + del exc_type, exc, traceback + self.release() + + +def reserve_loopback_port() -> LoopbackPortReservation: + """Return one live ephemeral reservation outside protected Comfy ports.""" for _ in range(16): - with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as probe: - probe.bind(("127.0.0.1", 0)) - port = int(probe.getsockname()[1]) + listener = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + try: + listener.bind(("127.0.0.1", 0)) + port = int(listener.getsockname()[1]) + except BaseException: + listener.close() + raise if port not in PROTECTED_COMFY_PORTS: - return port - raise RuntimeError("Unable to select a non-default loopback port.") + return LoopbackPortReservation(listener) + listener.close() + raise RuntimeError("Unable to reserve a non-default loopback port.") def validate_loopback_port(port: int) -> None: diff --git a/tools/comfy_integration/managed_server.py b/tools/comfy_integration/managed_server.py index e7c23a2..3715080 100644 --- a/tools/comfy_integration/managed_server.py +++ b/tools/comfy_integration/managed_server.py @@ -13,10 +13,12 @@ from pathlib import Path from tools.comfy_api import JsonObject, LoopbackComfyClient from .artifacts import IntegrationArtifacts -from .loopback_port import select_unused_loopback_port +from .loopback_port import is_loopback_port_available, reserve_loopback_port from .readiness import wait_for_server from .server_process import ComfyServerCommand, WindowsComfyProcess +MAX_LAUNCH_ATTEMPTS = 4 + @dataclass(frozen=True) class RunningManagedComfy: @@ -50,9 +52,22 @@ class ManagedComfyServer: self._running: RunningManagedComfy | None = None def __enter__(self) -> RunningManagedComfy: - """Start and verify a new loopback server.""" + """Start and verify a server, retrying only a proven port collision.""" + + for attempt in range(1, MAX_LAUNCH_ATTEMPTS + 1): + reservation = reserve_loopback_port() + port = reservation.port + reservation.release() + try: + return self._start_on_port(port) + except Exception: + if attempt == MAX_LAUNCH_ATTEMPTS or is_loopback_port_available(port): + raise + raise AssertionError("Managed Comfy launch attempts were not exhausted.") + + def _start_on_port(self, port: int) -> RunningManagedComfy: + """Launch and verify one server attempt on an owned candidate port.""" - port = select_unused_loopback_port() command = ComfyServerCommand( self._root, self._root / "venv" / "Scripts" / "python.exe", diff --git a/tools/negpip_integration/upstream_parity.py b/tools/negpip_integration/upstream_parity.py index 37af3d1..70b74ea 100644 --- a/tools/negpip_integration/upstream_parity.py +++ b/tools/negpip_integration/upstream_parity.py @@ -59,6 +59,7 @@ def prove_upstream_parity(ppm_root: Path) -> dict[str, object]: check=True, capture_output=True, text=True, + timeout=30.0, ).stdout.strip() if revision != PPM_REVISION: raise ValueError( diff --git a/tools/test_governance/__init__.py b/tools/test_governance/__init__.py new file mode 100644 index 0000000..89635cc --- /dev/null +++ b/tools/test_governance/__init__.py @@ -0,0 +1,5 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Enforce reviewed test architecture and execution constraints.""" diff --git a/tools/test_governance/ast_analysis.py b/tools/test_governance/ast_analysis.py new file mode 100644 index 0000000..8d5c893 --- /dev/null +++ b/tools/test_governance/ast_analysis.py @@ -0,0 +1,55 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Resolve imported Python call identities without executing test source.""" + +from __future__ import annotations + +import ast + + +def import_aliases(tree: ast.Module) -> dict[str, str]: + """Return local import names mapped to their canonical dotted owners.""" + + aliases: dict[str, str] = {} + for node in tree.body: + if isinstance(node, ast.Import): + for imported in node.names: + local_name = imported.asname or imported.name.split(".", maxsplit=1)[0] + aliases[local_name] = imported.name + elif isinstance(node, ast.ImportFrom) and node.module is not None: + for imported in node.names: + if imported.name == "*": + continue + local_name = imported.asname or imported.name + aliases[local_name] = f"{node.module}.{imported.name}" + return aliases + + +def call_name(node: ast.expr, aliases: dict[str, str]) -> str: + """Return one dotted call name without evaluating source.""" + + if isinstance(node, ast.Name): + return aliases.get(node.id, node.id) + if isinstance(node, ast.Attribute): + owner = call_name(node.value, aliases) + return f"{owner}.{node.attr}" if owner else node.attr + return "" + + +def configured_call_name( + call_identity: str, + configured_calls: frozenset[str], +) -> str | None: + """Return the most specific configured name matching one canonical call.""" + + matches = tuple( + configured + for configured in configured_calls + if call_identity == configured or call_identity.endswith(f".{configured}") + ) + return max(matches, key=len, default=None) + + +__all__ = ["call_name", "configured_call_name", "import_aliases"] diff --git a/tools/test_governance/discovery.py b/tools/test_governance/discovery.py new file mode 100644 index 0000000..63e4de2 --- /dev/null +++ b/tools/test_governance/discovery.py @@ -0,0 +1,347 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Orchestrate objective test-pattern discovery across focused rule owners.""" + +from __future__ import annotations + +import ast +import re +from pathlib import Path + +from .ast_analysis import import_aliases +from .execution_patterns import execution_pattern_candidates, reads_environment_name +from .model import TestCandidate, TestPolicy +from .network_resource_patterns import closed_ephemeral_port_candidates +from .node_process_patterns import node_process_pattern_candidates +from .ownership_patterns import ownership_pattern_candidates +from .process_lifecycle_patterns import process_lifecycle_pattern_candidates +from .process_state_patterns import process_state_pattern_candidates +from .semantic_patterns import semantic_pattern_candidates + +LAYOUT_RULE = "LAYOUT001" +STUB_RULE = "STUB001" +XDIST_RULE = "XDIST001" +SERIAL_RULE = "SERIAL001" +ISOLATED_RULE = "ISOLATED001" +SCRATCH_RULE = "SCRATCH001" +TYPESCRIPT_LAYOUT_RULE = "TSLAYOUT001" +TYPESCRIPT_OPTIONAL_RULE = "TSOPTIONAL001" + +_TYPESCRIPT_OPTIONAL_PATTERN = re.compile( + r"\b(?:describe|it|test)\.(?:only|skip|todo)\s*\(" +) + + +def discover_test_candidates( + root: Path, policy: TestPolicy +) -> tuple[TestCandidate, ...]: + """Return every exact current test-governance review candidate.""" + + test_root = root / policy.test_root + candidates = [ + *_root_layout_candidates(root, test_root, policy), + *_typescript_test_candidates(root), + *_stub_candidates(root, test_root), + *_execution_inventory_candidates( + root, + policy, + variable_name="ISOLATED_TEST_MODULES", + rule=ISOLATED_RULE, + locator="isolated-module", + evidence="module is assigned to the bounded fresh-process runner", + ), + *_execution_inventory_candidates( + root, + policy, + variable_name="SERIAL_TEST_MODULES", + rule=SERIAL_RULE, + locator="serial-module", + evidence="module is assigned to the globally sequential serial runner", + ), + ] + for path in sorted(test_root.rglob("*.py")): + if "__pycache__" in path.parts: + continue + candidates.extend(_python_source_candidates(root, path, policy)) + for support_root in policy.semantic_support_roots: + for path in sorted((root / support_root).rglob("*.py")): + if "__pycache__" in path.parts or path.is_relative_to(test_root): + continue + candidates.extend(_semantic_support_source_candidates(root, path)) + return tuple( + sorted(candidates, key=lambda item: (item.path, item.rule, item.locator)) + ) + + +def _typescript_test_candidates(root: Path) -> list[TestCandidate]: + """Find misplaced or optional frontend tests in the authored test tree.""" + + test_root = root / "web/tests" + if not test_root.is_dir(): + return [] + candidates = [ + TestCandidate( + rule=TYPESCRIPT_LAYOUT_RULE, + path=path.relative_to(root).as_posix(), + locator="frontend-root-source", + evidence="authored source is stored directly under web/tests/", + line=1, + ) + for path in sorted(test_root.iterdir()) + if path.is_file() and path.suffix == ".ts" + ] + for path in sorted(test_root.rglob("*.ts")): + for line_number, line in enumerate( + path.read_text(encoding="utf-8").splitlines(), + start=1, + ): + if _TYPESCRIPT_OPTIONAL_PATTERN.search(line): + candidates.append( + TestCandidate( + rule=TYPESCRIPT_OPTIONAL_RULE, + path=path.relative_to(root).as_posix(), + locator=f"optional-proof:{line_number}", + evidence="frontend proof uses an optional test declaration", + line=line_number, + ) + ) + return candidates + + +def _root_layout_candidates( + root: Path, + test_root: Path, + policy: TestPolicy, +) -> list[TestCandidate]: + """Find authored test sources still stored at the test-package root.""" + + candidates: list[TestCandidate] = [] + for path in sorted(test_root.iterdir()): + relative_path = path.relative_to(root).as_posix() + if ( + not path.is_file() + or path.suffix not in policy.root_source_extensions + or relative_path in policy.allowed_root_source_paths + ): + continue + candidates.append( + TestCandidate( + rule=LAYOUT_RULE, + path=relative_path, + locator="root-source", + evidence="authored source is stored directly under tests/", + line=1, + ) + ) + return candidates + + +def _stub_candidates(root: Path, test_root: Path) -> list[TestCandidate]: + """Find test stubs that can shadow executable test modules or typing state.""" + + return [ + TestCandidate( + rule=STUB_RULE, + path=path.relative_to(root).as_posix(), + locator="test-stub", + evidence=( + "stub shadows a Python module" + if path.with_suffix(".py").is_file() + else "standalone test stub bypasses executable strict typing" + ), + line=1, + ) + for path in sorted(test_root.rglob("*.pyi")) + if "__pycache__" not in path.parts + ] + + +def _execution_inventory_candidates( + root: Path, + policy: TestPolicy, + *, + variable_name: str, + rule: str, + locator: str, + evidence: str, +) -> list[TestCandidate]: + """Find every module assigned to one constrained execution inventory.""" + + policy_path = root / policy.serial_policy + tree = ast.parse(policy_path.read_text(encoding="utf-8"), filename=str(policy_path)) + module_paths: set[str] = set() + for node in tree.body: + if not isinstance(node, (ast.Assign, ast.AnnAssign)): + continue + targets = node.targets if isinstance(node, ast.Assign) else [node.target] + if not any( + isinstance(target, ast.Name) and target.id == variable_name + for target in targets + ): + continue + value = node.value + if value is None: + continue + module_paths.update( + constant.value + for constant in ast.walk(value) + if isinstance(constant, ast.Constant) + and isinstance(constant.value, str) + and constant.value.startswith("tests/") + ) + return [ + TestCandidate( + rule=rule, + path=module_path, + locator=locator, + evidence=evidence, + line=1, + ) + for module_path in sorted(module_paths) + ] + + +def _python_source_candidates( + root: Path, + path: Path, + policy: TestPolicy, +) -> list[TestCandidate]: + """Delegate one Python test source to its pattern owners.""" + + relative_path = path.relative_to(root).as_posix() + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + aliases = import_aliases(tree) + candidates: list[TestCandidate] = [] + xdist_reference = next( + ( + node + for node in ast.walk(tree) + if reads_environment_name( + node, + policy.xdist_environment_name, + aliases, + ) + ), + None, + ) + if xdist_reference is not None: + candidates.append( + TestCandidate( + rule=XDIST_RULE, + path=relative_path, + locator="module-xdist-branch", + evidence=f"source references {policy.xdist_environment_name}", + line=getattr(xdist_reference, "lineno", 1), + ) + ) + scratch_reference = next( + ( + node + for node in ast.walk(tree) + if isinstance(node, ast.Constant) + and isinstance(node.value, str) + and len(node.value) < 256 + and policy.repository_scratch_name in node.value + ), + None, + ) + if scratch_reference is not None: + candidates.append( + TestCandidate( + rule=SCRATCH_RULE, + path=relative_path, + locator="repository-scratch-reference", + evidence=f"source references {policy.repository_scratch_name}", + line=scratch_reference.lineno, + ) + ) + candidates.extend( + execution_pattern_candidates( + path=relative_path, + tree=tree, + wait_calls=policy.wait_calls, + wall_clock_calls=policy.wall_clock_calls, + aliases=aliases, + ) + ) + candidates.extend( + ownership_pattern_candidates( + root=root, + test_root=root / policy.test_root, + source_path=path, + relative_path=relative_path, + tree=tree, + aliases=aliases, + ) + ) + candidates.extend( + semantic_pattern_candidates( + relative_path=relative_path, + tree=tree, + aliases=aliases, + ) + ) + candidates.extend(closed_ephemeral_port_candidates(path=relative_path, tree=tree)) + candidates.extend( + process_state_pattern_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ) + ) + candidates.extend( + process_lifecycle_pattern_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ) + ) + candidates.extend( + node_process_pattern_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ) + ) + return candidates + + +def _semantic_support_source_candidates(root: Path, path: Path) -> list[TestCandidate]: + """Discover semantic reliability risks in test-owned support tooling.""" + + relative_path = path.relative_to(root).as_posix() + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + aliases = import_aliases(tree) + return [ + *semantic_pattern_candidates( + relative_path=relative_path, + tree=tree, + aliases=aliases, + ), + *process_lifecycle_pattern_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + *node_process_pattern_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + *closed_ephemeral_port_candidates(path=relative_path, tree=tree), + ] + + +__all__ = [ + "ISOLATED_RULE", + "LAYOUT_RULE", + "SCRATCH_RULE", + "SERIAL_RULE", + "STUB_RULE", + "TYPESCRIPT_LAYOUT_RULE", + "TYPESCRIPT_OPTIONAL_RULE", + "XDIST_RULE", + "discover_test_candidates", +] diff --git a/tools/test_governance/execution_patterns.py b/tools/test_governance/execution_patterns.py new file mode 100644 index 0000000..cb52654 --- /dev/null +++ b/tools/test_governance/execution_patterns.py @@ -0,0 +1,340 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover timed, polled, environmental, and fixed-resource test patterns.""" + +from __future__ import annotations + +import ast +from collections import Counter + +from .ast_analysis import call_name, configured_call_name +from .model import TestCandidate + +WAIT_RULE = "WAIT001" +CLOCK_RULE = "CLOCK001" +POLL_RULE = "POLL001" +ENVIRONMENT_RULE = "ENV001" +RESOURCE_RULE = "RESOURCE001" + + +def execution_pattern_candidates( + *, + path: str, + tree: ast.Module, + wait_calls: frozenset[str], + wall_clock_calls: frozenset[str], + aliases: dict[str, str], +) -> list[TestCandidate]: + """Return time, polling, environment, and resource candidates.""" + + visitor = _ExecutionPatternVisitor( + path=path, + wait_calls=wait_calls, + wall_clock_calls=wall_clock_calls, + aliases=aliases, + ) + visitor.visit(tree) + return [*visitor.candidates, *visitor.wall_clock_candidates()] + + +def reads_environment_name( + node: ast.AST, + environment_name: str, + aliases: dict[str, str], +) -> bool: + """Return whether one expression reads an exact environment variable.""" + + if isinstance(node, ast.Call) and call_name(node.func, aliases) in { + "os.environ.get", + "os.getenv", + }: + return bool( + node.args + and isinstance(node.args[0], ast.Constant) + and node.args[0].value == environment_name + ) + if ( + isinstance(node, ast.Subscript) + and call_name(node.value, aliases) == "os.environ" + ): + return ( + isinstance(node.slice, ast.Constant) + and node.slice.value == environment_name + ) + return False + + +class _ExecutionPatternVisitor(ast.NodeVisitor): + """Collect stable scoped locators for time and resource patterns.""" + + def __init__( + self, + *, + path: str, + wait_calls: frozenset[str], + wall_clock_calls: frozenset[str], + aliases: dict[str, str], + ) -> None: + """Initialize one source visitor with exact configured call names.""" + + self._path = path + self._wait_calls = wait_calls + self._wall_clock_calls = wall_clock_calls + self._aliases = aliases + self._scope: list[str] = [""] + self._counts: Counter[tuple[str, str]] = Counter() + self._clock_scopes: set[str] = set() + self._asserted_comparisons: list[tuple[str, ast.Compare]] = [] + self.candidates: list[TestCandidate] = [] + + def visit_ClassDef(self, node: ast.ClassDef) -> None: + """Track class ownership while visiting its body.""" + + self._scope.append(node.name) + self.generic_visit(node) + self._scope.pop() + + def visit_FunctionDef(self, node: ast.FunctionDef) -> None: + """Track function ownership while visiting its body.""" + + self._scope.append(node.name) + self.generic_visit(node) + self._scope.pop() + + def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> None: + """Track async-function ownership while visiting its body.""" + + self._scope.append(node.name) + self.generic_visit(node) + self._scope.pop() + + def visit_Call(self, node: ast.Call) -> None: + """Record configured waits, clock use, and fixed socket binds.""" + + identity = call_name(node.func, self._aliases) + wait_name = configured_call_name(identity, self._wait_calls) + scope = self._scope_name + if configured_call_name(identity, self._wall_clock_calls) is not None: + self._clock_scopes.add(scope) + if wait_name is not None: + ordinal = self._next_ordinal(scope, f"wait-{wait_name}") + self.candidates.append( + TestCandidate( + rule=WAIT_RULE, + path=self._path, + locator=f"{scope}:wait:{wait_name}:{ordinal}", + evidence=f"calls {wait_name}", + line=node.lineno, + ) + ) + if identity.endswith(".bind") and _has_fixed_bind_port(node): + ordinal = self._next_ordinal(scope, "fixed-bind") + self.candidates.append( + TestCandidate( + rule=RESOURCE_RULE, + path=self._path, + locator=f"{scope}:fixed-bind:{ordinal}", + evidence="binds a socket to a fixed nonzero port", + line=node.lineno, + ) + ) + if identity in { + "os.environ.clear", + "os.environ.pop", + "os.environ.popitem", + "os.environ.setdefault", + "os.environ.update", + }: + self._record_environment_mutation(node) + self.generic_visit(node) + + def visit_Assign(self, node: ast.Assign) -> None: + """Record direct writes to the process environment.""" + + if any(self._is_environment_target(target) for target in node.targets): + self._record_environment_mutation(node) + self.generic_visit(node) + + def visit_AnnAssign(self, node: ast.AnnAssign) -> None: + """Record annotated writes to the process environment.""" + + if self._is_environment_target(node.target): + self._record_environment_mutation(node) + self.generic_visit(node) + + def visit_AugAssign(self, node: ast.AugAssign) -> None: + """Record augmented writes to the process environment.""" + + if self._is_environment_target(node.target): + self._record_environment_mutation(node) + self.generic_visit(node) + + def visit_Delete(self, node: ast.Delete) -> None: + """Record direct deletion from the process environment.""" + + if any(self._is_environment_target(target) for target in node.targets): + self._record_environment_mutation(node) + self.generic_visit(node) + + def visit_While(self, node: ast.While) -> None: + """Record loops whose completion or failure bound reads a real clock.""" + + scope = self._scope_name + direct_clock = any( + isinstance(item, ast.Call) + and configured_call_name( + call_name(item.func, self._aliases), + self._wall_clock_calls, + ) + is not None + for item in ast.walk(node) + ) + elapsed_timer = scope in self._clock_scopes and any( + isinstance(item, ast.Call) + and isinstance(item.func, ast.Attribute) + and item.func.attr == "elapsed" + for item in ast.walk(node) + ) + if direct_clock or elapsed_timer: + ordinal = self._next_ordinal(scope, "wall-clock-poll") + self.candidates.append( + TestCandidate( + rule=POLL_RULE, + path=self._path, + locator=f"{scope}:wall-clock-poll:{ordinal}", + evidence="bounds a polling loop with a real wall clock", + line=node.lineno, + ) + ) + self.generic_visit(node) + + def visit_Assert(self, node: ast.Assert) -> None: + """Retain asserted comparisons for clock-scope analysis after traversal.""" + + self._asserted_comparisons.extend( + (self._scope_name, comparison) + for comparison in ast.walk(node.test) + if isinstance(comparison, ast.Compare) + ) + self.generic_visit(node) + + def wall_clock_candidates(self) -> list[TestCandidate]: + """Return numeric timing thresholds from scopes that read a real clock.""" + + candidates: list[TestCandidate] = [] + for scope, comparison in self._asserted_comparisons: + if scope not in self._clock_scopes or not _is_timing_threshold( + comparison, + self._wall_clock_calls, + self._aliases, + ): + continue + ordinal = self._next_ordinal(scope, "wall-clock-threshold") + candidates.append( + TestCandidate( + rule=CLOCK_RULE, + path=self._path, + locator=f"{scope}:wall-clock-threshold:{ordinal}", + evidence="compares real elapsed time with a numeric threshold", + line=comparison.lineno, + ) + ) + return candidates + + @property + def _scope_name(self) -> str: + """Return the stable qualified scope currently being visited.""" + + return ".".join(self._scope) + + def _next_ordinal(self, scope: str, pattern: str) -> int: + """Return a stable one-based occurrence ordinal within one scope.""" + + key = (scope, pattern) + self._counts[key] += 1 + return self._counts[key] + + def _is_environment_target(self, node: ast.expr) -> bool: + """Return whether one assignment target belongs to ``os.environ``.""" + + return ( + isinstance(node, ast.Subscript) + and call_name(node.value, self._aliases) == "os.environ" + ) + + def _record_environment_mutation(self, node: ast.AST) -> None: + """Record one exact mutation of process-global environment state.""" + + scope = self._scope_name + ordinal = self._next_ordinal(scope, "environment-mutation") + self.candidates.append( + TestCandidate( + rule=ENVIRONMENT_RULE, + path=self._path, + locator=f"{scope}:environment-mutation:{ordinal}", + evidence="mutates process-global environment state directly", + line=getattr(node, "lineno", 1), + ) + ) + + +def _has_fixed_bind_port(call: ast.Call) -> bool: + """Return whether one socket bind call contains a fixed nonzero port.""" + + if not call.args or not isinstance(call.args[0], (ast.Tuple, ast.List)): + return False + elements = call.args[0].elts + return ( + len(elements) >= 2 + and isinstance(elements[1], ast.Constant) + and isinstance(elements[1].value, int) + and elements[1].value > 0 + ) + + +def _is_timing_threshold( + comparison: ast.Compare, + wall_clock_calls: frozenset[str], + aliases: dict[str, str], +) -> bool: + """Return whether a comparison imposes a numeric real-time threshold.""" + + if not any( + isinstance(operator, (ast.Lt, ast.LtE, ast.Gt, ast.GtE)) + for operator in comparison.ops + ): + return False + expressions = [comparison.left, *comparison.comparators] + has_number = any( + isinstance(expression, ast.Constant) + and isinstance(expression.value, (int, float)) + and not isinstance(expression.value, bool) + for expression in expressions + ) + names = { + node.id.casefold() + for expression in expressions + for node in ast.walk(expression) + if isinstance(node, ast.Name) + } + timing_name = any( + token in name or name.endswith(("_ms", "_seconds")) + for name in names + for token in ("elapsed", "duration", "latency", "timeout", "deadline") + ) + reads_clock = any( + isinstance(node, ast.Call) + and configured_call_name( + call_name(node.func, aliases), + wall_clock_calls, + ) + is not None + for expression in expressions + for node in ast.walk(expression) + ) + return has_number and (timing_name or reads_clock) + + +__all__ = ["execution_pattern_candidates", "reads_environment_name"] diff --git a/tools/test_governance/loading.py b/tools/test_governance/loading.py new file mode 100644 index 0000000..0dddc30 --- /dev/null +++ b/tools/test_governance/loading.py @@ -0,0 +1,225 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Load strict test policy, debt, and waiver state.""" + +from __future__ import annotations + +import tomllib +from collections.abc import Mapping +from datetime import date +from pathlib import Path +from typing import cast + +from .model import TestDebt, TestPolicy, TestState, TestWaiver + + +def load_test_policy(path: Path) -> TestPolicy: + """Load one schema-versioned test-governance policy.""" + + data = _document(path, expected_version=1) + _exact_keys(data, {"schema_version", "scope", "discovery", "registries"}, str(path)) + scope = _mapping(data, "scope") + discovery = _mapping(data, "discovery") + registries = _mapping(data, "registries") + _exact_keys( + scope, + { + "test_root", + "semantic_support_roots", + "root_source_extensions", + "allowed_root_source_paths", + }, + "test scope", + ) + _exact_keys( + discovery, + { + "serial_policy", + "wait_calls", + "wall_clock_calls", + "xdist_environment_name", + "repository_scratch_name", + }, + "test discovery", + ) + _exact_keys(registries, {"debt", "waivers"}, "test registries") + return TestPolicy( + test_root=Path(_string(scope, "test_root")), + semantic_support_roots=tuple( + Path(value) for value in _strings(scope, "semantic_support_roots") + ), + root_source_extensions=frozenset(_strings(scope, "root_source_extensions")), + allowed_root_source_paths=frozenset( + _strings(scope, "allowed_root_source_paths") + ), + serial_policy=Path(_string(discovery, "serial_policy")), + wait_calls=frozenset(_strings(discovery, "wait_calls")), + wall_clock_calls=frozenset(_strings(discovery, "wall_clock_calls")), + xdist_environment_name=_string(discovery, "xdist_environment_name"), + repository_scratch_name=_string(discovery, "repository_scratch_name"), + debt_registry=Path(_string(registries, "debt")), + waiver_registry=Path(_string(registries, "waivers")), + ) + + +def load_test_state(root: Path, policy: TestPolicy) -> TestState: + """Load exact current test debt and waiver registries.""" + + debt_data = _registry_document(root / policy.debt_registry, "debts") + waiver_data = _registry_document(root / policy.waiver_registry, "waivers") + return TestState( + debts=tuple(_parse_debt(item) for item in _tables(debt_data, "debts")), + waivers=tuple(_parse_waiver(item) for item in _tables(waiver_data, "waivers")), + ) + + +def _parse_debt(data: Mapping[str, object]) -> TestDebt: + """Parse one exact test-debt record.""" + + _exact_keys( + data, + { + "id", + "owner", + "rule", + "candidates", + "paths", + "fingerprint", + "issue", + "review_by", + "problem", + "remediation", + }, + "test debt", + ) + return TestDebt( + identifier=_string(data, "id"), + owner=_string(data, "owner"), + rule=_string(data, "rule"), + candidates=_strings(data, "candidates"), + paths=_strings(data, "paths"), + fingerprint=_string(data, "fingerprint"), + issue=_string(data, "issue"), + review_by=_date(data, "review_by"), + problem=_string(data, "problem"), + remediation=_string(data, "remediation"), + ) + + +def _parse_waiver(data: Mapping[str, object]) -> TestWaiver: + """Parse one exact reviewed test waiver.""" + + kind = _string(data, "kind") + if kind not in {"classification", "remediation"}: + raise ValueError("test waiver kind must be classification or remediation") + fields = { + "id", + "owner", + "kind", + "disposition", + "rule", + "candidates", + "paths", + "fingerprint", + "rationale", + "issue", + "review_by", + } + if kind == "remediation": + fields.add("debt") + _exact_keys(data, fields, "test waiver") + return TestWaiver( + identifier=_string(data, "id"), + owner=_string(data, "owner"), + kind=kind, + disposition=_string(data, "disposition"), + rule=_string(data, "rule"), + candidates=_strings(data, "candidates"), + paths=_strings(data, "paths"), + fingerprint=_string(data, "fingerprint"), + rationale=_string(data, "rationale"), + issue=_string(data, "issue"), + review_by=_date(data, "review_by"), + debt=_string(data, "debt") if kind == "remediation" else None, + ) + + +def _document(path: Path, *, expected_version: int) -> Mapping[str, object]: + """Read one TOML document and validate its schema version.""" + + data = cast(Mapping[str, object], tomllib.loads(path.read_text(encoding="utf-8"))) + if data.get("schema_version") != expected_version: + raise ValueError(f"{path} must declare schema_version = {expected_version}") + return data + + +def _registry_document(path: Path, collection: str) -> Mapping[str, object]: + """Read one strict current-state registry envelope.""" + + data = _document(path, expected_version=1) + _exact_keys(data, {"schema_version", collection}, str(path)) + return data + + +def _exact_keys(data: Mapping[str, object], expected: set[str], label: str) -> None: + """Reject missing and history-shaped surplus fields.""" + + missing = expected - set(data) + unknown = set(data) - expected + if missing or unknown: + raise ValueError( + f"{label} fields differ: missing={sorted(missing)}, " + f"unsupported={sorted(unknown)}" + ) + + +def _mapping(data: Mapping[str, object], key: str) -> Mapping[str, object]: + """Return one required TOML table.""" + + value = data.get(key) + if not isinstance(value, dict): + raise TypeError(f"{key} must be a table") + return cast(Mapping[str, object], value) + + +def _tables(data: Mapping[str, object], key: str) -> tuple[Mapping[str, object], ...]: + """Return one optional array of TOML tables.""" + + value = data.get(key, []) + if not isinstance(value, list) or not all(isinstance(item, dict) for item in value): + raise TypeError(f"{key} must be an array of tables") + return tuple(cast(list[Mapping[str, object]], value)) + + +def _string(data: Mapping[str, object], key: str) -> str: + """Return one required nonempty string.""" + + value = data.get(key) + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{key} must be a nonempty string") + return value + + +def _strings(data: Mapping[str, object], key: str) -> tuple[str, ...]: + """Return one required nonempty array of unique strings.""" + + value = data.get(key) + if not isinstance(value, list) or not value: + raise ValueError(f"{key} must be a nonempty string array") + if not all(isinstance(item, str) and item.strip() for item in value): + raise ValueError(f"{key} must contain nonempty strings") + strings = tuple(cast(list[str], value)) + if len(strings) != len(set(strings)): + raise ValueError(f"{key} must not contain duplicates") + return strings + + +def _date(data: Mapping[str, object], key: str) -> date: + """Return one required TOML date.""" + + value = data.get(key) + if not isinstance(value, date): + raise TypeError(f"{key} must be an ISO date") + return value diff --git a/tools/test_governance/metrics.py b/tools/test_governance/metrics.py new file mode 100644 index 0000000..99b95be --- /dev/null +++ b/tools/test_governance/metrics.py @@ -0,0 +1,32 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own stable fingerprints for reviewed test-governance facts.""" + +from __future__ import annotations + +import hashlib +from pathlib import Path + +from tools.architecture_governance.metrics import source_fingerprint + +_INVENTORY_RULES = frozenset({"ISOLATED001", "LAYOUT001", "SERIAL001", "STUB001"}) + + +def reviewed_state_fingerprint( + root: Path, + *, + rule: str, + candidates: tuple[str, ...], + paths: tuple[str, ...], +) -> str: + """Fingerprint the exact fact whose reviewed disposition must remain stable.""" + + if rule not in _INVENTORY_RULES: + return source_fingerprint(root, paths) + digest = hashlib.sha256() + for candidate in sorted(candidates): + digest.update(candidate.encode("utf-8")) + digest.update(b"\0") + return f"sha256:{digest.hexdigest()}" diff --git a/tools/test_governance/model.py b/tools/test_governance/model.py new file mode 100644 index 0000000..16f1411 --- /dev/null +++ b/tools/test_governance/model.py @@ -0,0 +1,97 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Define immutable test-governance policy and reviewed state.""" + +from __future__ import annotations + +from dataclasses import dataclass +from datetime import date +from pathlib import Path + +from tools.architecture_governance.model import Diagnostic + + +@dataclass(frozen=True, slots=True) +class TestPolicy: + """Define exact test discovery inputs and registry locations.""" + + test_root: Path + semantic_support_roots: tuple[Path, ...] + root_source_extensions: frozenset[str] + allowed_root_source_paths: frozenset[str] + serial_policy: Path + wait_calls: frozenset[str] + wall_clock_calls: frozenset[str] + xdist_environment_name: str + repository_scratch_name: str + debt_registry: Path + waiver_registry: Path + + +@dataclass(frozen=True, slots=True) +class TestCandidate: + """Identify one mechanically discovered fact requiring human review.""" + + rule: str + path: str + locator: str + evidence: str + line: int + + @property + def key(self) -> str: + """Return the stable identity used by reviewed state records.""" + + return f"{self.rule}|{self.path}|{self.locator}" + + +@dataclass(frozen=True, slots=True) +class TestDebt: + """Describe reviewed test debt and its concrete remediation.""" + + identifier: str + owner: str + rule: str + candidates: tuple[str, ...] + paths: tuple[str, ...] + fingerprint: str + issue: str + review_by: date + problem: str + remediation: str + + +@dataclass(frozen=True, slots=True) +class TestWaiver: + """Describe an exact classification or debt-remediation exception.""" + + identifier: str + owner: str + kind: str + disposition: str + rule: str + candidates: tuple[str, ...] + paths: tuple[str, ...] + fingerprint: str + rationale: str + issue: str + review_by: date + debt: str | None + + +@dataclass(frozen=True, slots=True) +class TestState: + """Collect current test debt and waiver snapshots.""" + + debts: tuple[TestDebt, ...] + waivers: tuple[TestWaiver, ...] + + +@dataclass(frozen=True, slots=True) +class TestValidationResult: + """Return discovered candidates together with policy diagnostics.""" + + candidates: tuple[TestCandidate, ...] + diagnostics: tuple[Diagnostic, ...] diff --git a/tools/test_governance/network_resource_patterns.py b/tools/test_governance/network_resource_patterns.py new file mode 100644 index 0000000..cb167b6 --- /dev/null +++ b/tools/test_governance/network_resource_patterns.py @@ -0,0 +1,134 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover network-resource lifetime risks in test and CI support.""" + +from __future__ import annotations + +import ast + +from .model import TestCandidate + +PORT_HANDOFF_RULE = "PORT001" + + +def closed_ephemeral_port_candidates( + *, + path: str, + tree: ast.Module, +) -> list[TestCandidate]: + """Find port helpers that close their reservation before returning its number.""" + + candidates: list[TestCandidate] = [] + ordinal = 0 + for function in ( + node + for node in ast.walk(tree) + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + ): + for context in ( + node for node in ast.walk(function) if isinstance(node, ast.With) + ): + for item in context.items: + if not isinstance(item.optional_vars, ast.Name): + continue + socket_name = item.optional_vars.id + if not _binds_os_assigned_port(context, socket_name): + continue + derived_names = _port_names_derived_from_socket(context, socket_name) + if not any( + _returns_socket_port(statement, socket_name, derived_names) + for statement in function.body + ): + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=PORT_HANDOFF_RULE, + path=path, + locator=f"{function.name}:closed-port-handoff:{ordinal}", + evidence=( + "returns an OS-assigned port after its reservation socket " + "closes" + ), + line=context.lineno, + ) + ) + return candidates + + +def _binds_os_assigned_port(context: ast.With, socket_name: str) -> bool: + """Return whether one context asks the OS for a port on its owned socket.""" + + return any( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and isinstance(node.func.value, ast.Name) + and node.func.value.id == socket_name + and node.func.attr == "bind" + and any( + isinstance(value, ast.Constant) and value.value == 0 + for argument in node.args + for value in ast.walk(argument) + ) + for statement in context.body + for node in ast.walk(statement) + ) + + +def _port_names_derived_from_socket( + context: ast.With, + socket_name: str, +) -> frozenset[str]: + """Return local names assigned from the owned socket's address.""" + + return frozenset( + target.id + for statement in context.body + for node in ast.walk(statement) + if isinstance(node, (ast.Assign, ast.AnnAssign)) + for target in ( + (*node.targets,) if isinstance(node, ast.Assign) else (node.target,) + ) + if isinstance(target, ast.Name) + and node.value is not None + and _calls_socket_getsockname(node.value, socket_name) + ) + + +def _returns_socket_port( + statement: ast.stmt, + socket_name: str, + derived_names: frozenset[str], +) -> bool: + """Return whether a statement publishes a closed socket's derived port.""" + + return any( + isinstance(node, ast.Return) + and node.value is not None + and ( + _calls_socket_getsockname(node.value, socket_name) + or any( + isinstance(value, ast.Name) and value.id in derived_names + for value in ast.walk(node.value) + ) + ) + for node in ast.walk(statement) + ) + + +def _calls_socket_getsockname(expression: ast.expr, socket_name: str) -> bool: + """Return whether an expression reads the owned socket address.""" + + return any( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and isinstance(node.func.value, ast.Name) + and node.func.value.id == socket_name + and node.func.attr == "getsockname" + for node in ast.walk(expression) + ) + + +__all__ = ["PORT_HANDOFF_RULE", "closed_ephemeral_port_candidates"] diff --git a/tools/test_governance/node_process_patterns.py b/tools/test_governance/node_process_patterns.py new file mode 100644 index 0000000..1832b9d --- /dev/null +++ b/tools/test_governance/node_process_patterns.py @@ -0,0 +1,148 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover real Node commands that bypass bounded test execution ownership.""" + +from __future__ import annotations + +import ast + +from .ast_analysis import call_name +from .model import TestCandidate + +NODE_PROCESS_RULE = "NODE001" + +_NODE_PROCESS_CALLS = frozenset( + { + "subprocess.Popen", + "subprocess.call", + "subprocess.check_call", + "subprocess.check_output", + "subprocess.run", + } +) +_LEXICAL_SCOPES = ( + ast.Module, + ast.FunctionDef, + ast.AsyncFunctionDef, + ast.Lambda, + ast.ClassDef, +) + + +def node_process_pattern_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find real Node commands that bypass the bounded test execution owner.""" + + parent_by_node = { + child: parent + for parent in ast.walk(tree) + for child in ast.iter_child_nodes(parent) + } + assignments = _unique_assignment_values_by_scope(tree, parent_by_node) + candidates: list[TestCandidate] = [] + ordinal = 0 + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + identity = call_name(node.func, aliases) + if identity not in _NODE_PROCESS_CALLS: + continue + command = _call_command_expression(node) + if isinstance(command, ast.Name): + scope = _lexical_scope(node, parent_by_node) + command = assignments.get(scope, {}).get(command.id) + if not _literal_command_starts_with_node(command): + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=NODE_PROCESS_RULE, + path=path, + locator=f":unowned-node-process:{ordinal}", + evidence=( + f"calls real Node through {identity} instead of the bounded " + "shared test runtime owner" + ), + line=node.lineno, + ) + ) + return candidates + + +def _unique_assignment_values_by_scope( + tree: ast.Module, + parent_by_node: dict[ast.AST, ast.AST], +) -> dict[ast.AST, dict[str, ast.expr]]: + """Return unambiguous simple-name values within their lexical scopes.""" + + bindings: dict[ast.AST, dict[str, list[ast.expr]]] = {} + for node in ast.walk(tree): + if isinstance(node, ast.Assign): + targets = node.targets + value = node.value + elif isinstance(node, ast.AnnAssign) and node.value is not None: + targets = [node.target] + value = node.value + else: + continue + scope = _lexical_scope(node, parent_by_node) + scope_bindings = bindings.setdefault(scope, {}) + for target in targets: + if isinstance(target, ast.Name): + scope_bindings.setdefault(target.id, []).append(value) + return { + scope: { + name: values[0] + for name, values in scope_bindings.items() + if len(values) == 1 + } + for scope, scope_bindings in bindings.items() + } + + +def _lexical_scope( + node: ast.AST, + parent_by_node: dict[ast.AST, ast.AST], +) -> ast.AST: + """Return the nearest scope that owns one expression or assignment.""" + + current = node + while current in parent_by_node: + current = parent_by_node[current] + if isinstance(current, _LEXICAL_SCOPES): + return current + return current + + +def _call_command_expression(node: ast.Call) -> ast.expr | None: + """Return the command expression from a subprocess call.""" + + if node.args: + return node.args[0] + return next( + (keyword.value for keyword in node.keywords if keyword.arg == "args"), + None, + ) + + +def _literal_command_starts_with_node(node: ast.expr | None) -> bool: + """Return whether one literal command invokes the Node executable.""" + + if not isinstance(node, (ast.List, ast.Tuple)) or not node.elts: + return False + executable = node.elts[0] + if not isinstance(executable, ast.Constant) or not isinstance( + executable.value, str + ): + return False + basename = executable.value.replace("\\", "/").rsplit("/", maxsplit=1)[-1] + return basename.casefold() in {"node", "node.exe"} + + +__all__ = ["NODE_PROCESS_RULE", "node_process_pattern_candidates"] diff --git a/tools/test_governance/ownership_patterns.py b/tools/test_governance/ownership_patterns.py new file mode 100644 index 0000000..ec56c47 --- /dev/null +++ b/tools/test_governance/ownership_patterns.py @@ -0,0 +1,171 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover executable-test coupling and import-time resource ownership risks.""" + +from __future__ import annotations + +import ast +from pathlib import Path + +from .ast_analysis import call_name +from .model import TestCandidate + +SIBLING_IMPORT_RULE = "IMPORT001" +MODULE_RESOURCE_RULE = "SCOPE001" + +_MUTABLE_RESOURCE_CONSTRUCTORS = frozenset( + { + "concurrent.futures.ProcessPoolExecutor", + "concurrent.futures.ThreadPoolExecutor", + "http.server.HTTPServer", + "http.server.ThreadingHTTPServer", + "PySide6.QtCore.QProcess", + "PySide6.QtCore.QThread", + "PySide6.QtCore.QTimer", + "PySide6.QtNetwork.QNetworkAccessManager", + "PySide6.QtWidgets.QApplication", + "PySide6.QtWidgets.QWidget", + "socket.socket", + "subprocess.Popen", + "tempfile.TemporaryDirectory", + "threading.Thread", + } +) + + +def ownership_pattern_candidates( + *, + root: Path, + test_root: Path, + source_path: Path, + relative_path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Return coupling and import-time resource candidates for one source.""" + + return [ + *_sibling_test_module_import_candidates( + root=root, + test_root=test_root, + source_path=source_path, + relative_path=relative_path, + tree=tree, + ), + *_module_resource_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + ] + + +def _module_resource_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find mutable process or operating-system resources created at import time.""" + + candidates: list[TestCandidate] = [] + ordinal = 0 + for statement in tree.body: + if not isinstance(statement, (ast.Assign, ast.AnnAssign)): + continue + value = statement.value + if not isinstance(value, ast.Call): + continue + constructor = call_name(value.func, aliases) + if constructor not in _MUTABLE_RESOURCE_CONSTRUCTORS: + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=MODULE_RESOURCE_RULE, + path=path, + locator=f":mutable-resource:{ordinal}", + evidence=f"constructs {constructor} during module import", + line=statement.lineno, + ) + ) + return candidates + + +def _sibling_test_module_import_candidates( + *, + root: Path, + test_root: Path, + source_path: Path, + relative_path: str, + tree: ast.Module, +) -> list[TestCandidate]: + """Find imports whose resolved owner is another executable test module.""" + + candidates: list[TestCandidate] = [] + ordinal = 0 + for node in ast.walk(tree): + if not isinstance(node, (ast.Import, ast.ImportFrom)): + continue + imported_paths = _resolved_imported_python_paths(root, source_path, node) + test_modules = sorted( + { + imported_path + for imported_path in imported_paths + if imported_path != source_path + and imported_path.is_relative_to(test_root) + and imported_path.is_file() + and imported_path.name.startswith("test_") + } + ) + for imported_path in test_modules: + ordinal += 1 + imported_relative = imported_path.relative_to(root).as_posix() + candidates.append( + TestCandidate( + rule=SIBLING_IMPORT_RULE, + path=relative_path, + locator=f":test-module-import:{ordinal}", + evidence=f"imports executable test module {imported_relative}", + line=node.lineno, + ) + ) + return candidates + + +def _resolved_imported_python_paths( + root: Path, + source_path: Path, + node: ast.Import | ast.ImportFrom, +) -> tuple[Path, ...]: + """Resolve imported repository Python modules without importing source.""" + + if isinstance(node, ast.Import): + return tuple( + root.joinpath(*imported.name.split(".")).with_suffix(".py") + for imported in node.names + ) + if node.level: + source_package = source_path.relative_to(root).parent.parts + keep_count = max(0, len(source_package) - (node.level - 1)) + base_parts = source_package[:keep_count] + else: + base_parts = () + module_parts = () if node.module is None else tuple(node.module.split(".")) + module_path = root.joinpath(*base_parts, *module_parts) + imported_paths = [module_path.with_suffix(".py")] + imported_paths.extend( + module_path.joinpath(imported.name).with_suffix(".py") + for imported in node.names + if imported.name != "*" + ) + return tuple(imported_paths) + + +__all__ = [ + "MODULE_RESOURCE_RULE", + "SIBLING_IMPORT_RULE", + "ownership_pattern_candidates", +] diff --git a/tools/test_governance/process_lifecycle_patterns.py b/tools/test_governance/process_lifecycle_patterns.py new file mode 100644 index 0000000..f17ad81 --- /dev/null +++ b/tools/test_governance/process_lifecycle_patterns.py @@ -0,0 +1,58 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover child-process lifetime ownership risks in tests.""" + +from __future__ import annotations + +import ast + +from .ast_analysis import call_name +from .model import TestCandidate + +CHILD_PROCESS_RULE = "PROCESS001" + + +def process_lifecycle_pattern_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find child processes whose lifetime is not owned by a context manager.""" + + parent_by_node = { + child: parent + for parent in ast.walk(tree) + for child in ast.iter_child_nodes(parent) + } + candidates: list[TestCandidate] = [] + for node in ast.walk(tree): + if not ( + isinstance(node, ast.Call) + and call_name(node.func, aliases) == "subprocess.Popen" + ): + continue + parent = parent_by_node.get(node) + context_managed = ( + isinstance(parent, ast.withitem) and parent.context_expr is node + ) + if context_managed: + continue + candidates.append( + TestCandidate( + rule=CHILD_PROCESS_RULE, + path=path, + locator=f":unscoped-child-process:{len(candidates) + 1}", + evidence=( + "starts a child process without a context-managed lifetime; " + "termination and bounded cleanup require source review" + ), + line=node.lineno, + ) + ) + return candidates + + +__all__ = ["CHILD_PROCESS_RULE", "process_lifecycle_pattern_candidates"] diff --git a/tools/test_governance/process_state_patterns.py b/tools/test_governance/process_state_patterns.py new file mode 100644 index 0000000..02de598 --- /dev/null +++ b/tools/test_governance/process_state_patterns.py @@ -0,0 +1,167 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover unowned mutation of interpreter and Qt process-global state.""" + +from __future__ import annotations + +import ast + +from .ast_analysis import call_name +from .model import TestCandidate + +MODULE_REGISTRY_RULE = "MODULES001" +QT_GLOBAL_RULE = "QTGLOBAL001" +CURRENT_DIRECTORY_RULE = "CWD001" + +_QFLUENT_STATE_OWNER = "tests/presentation/theme/support.py" +_QFLUENT_MUTATIONS = frozenset( + { + "qfluentwidgets.setTheme", + "qfluentwidgets.setThemeColor", + } +) + + +def process_state_pattern_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Return process-global state mutation candidates for one test source.""" + + return [ + *_module_registry_mutation_candidates(path=path, tree=tree, aliases=aliases), + *_qfluent_global_mutation_candidates(path=path, tree=tree, aliases=aliases), + *_current_directory_mutation_candidates(path=path, tree=tree, aliases=aliases), + ] + + +def _current_directory_mutation_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find direct mutation of the process-global working directory.""" + + mutations = [ + node + for node in ast.walk(tree) + if isinstance(node, ast.Call) and call_name(node.func, aliases) == "os.chdir" + ] + return [ + TestCandidate( + rule=CURRENT_DIRECTORY_RULE, + path=path, + locator=f":current-directory-mutation:{ordinal}", + evidence="mutates the process-global current directory directly", + line=node.lineno, + ) + for ordinal, node in enumerate( + sorted(mutations, key=lambda item: item.lineno), + 1, + ) + ] + + +def _module_registry_mutation_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find unscoped mutation of Python's process-global module registry.""" + + mutations: list[ast.stmt | ast.Call] = [] + destructive_calls = { + "sys.modules.clear", + "sys.modules.pop", + "sys.modules.popitem", + "sys.modules.setdefault", + "sys.modules.update", + } + for node in ast.walk(tree): + if ( + isinstance(node, ast.Call) + and call_name(node.func, aliases) in destructive_calls + ): + mutations.append(node) + continue + targets: tuple[ast.expr, ...] + if isinstance(node, ast.Assign): + targets = tuple(node.targets) + elif isinstance(node, (ast.AnnAssign, ast.AugAssign)): + targets = (node.target,) + elif isinstance(node, ast.Delete): + targets = tuple(node.targets) + else: + continue + if any(_is_module_registry_target(target, aliases) for target in targets): + mutations.append(node) + return [ + TestCandidate( + rule=MODULE_REGISTRY_RULE, + path=path, + locator=f":module-registry-mutation:{ordinal}", + evidence="mutates process-global sys.modules state directly", + line=node.lineno, + ) + for ordinal, node in enumerate( + sorted(mutations, key=lambda item: item.lineno), + 1, + ) + ] + + +def _is_module_registry_target( + node: ast.expr, + aliases: dict[str, str], +) -> bool: + """Return whether one assignment or deletion target belongs to sys.modules.""" + + return ( + isinstance(node, ast.Subscript) + and call_name(node.value, aliases) == "sys.modules" + ) + + +def _qfluent_global_mutation_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find direct QFluent appearance mutation outside its restoration owner.""" + + if path == _QFLUENT_STATE_OWNER: + return [] + mutations = [ + (node, identity) + for node in ast.walk(tree) + if isinstance(node, ast.Call) + and (identity := call_name(node.func, aliases)) in _QFLUENT_MUTATIONS + ] + return [ + TestCandidate( + rule=QT_GLOBAL_RULE, + path=path, + locator=f":qfluent-global-mutation:{ordinal}", + evidence=f"calls {identity} outside the QFluent state owner", + line=node.lineno, + ) + for ordinal, (node, identity) in enumerate( + sorted(mutations, key=lambda item: item[0].lineno), + 1, + ) + ] + + +__all__ = [ + "CURRENT_DIRECTORY_RULE", + "MODULE_REGISTRY_RULE", + "QT_GLOBAL_RULE", + "process_state_pattern_candidates", +] diff --git a/tools/test_governance/semantic_patterns.py b/tools/test_governance/semantic_patterns.py new file mode 100644 index 0000000..4349d8a --- /dev/null +++ b/tools/test_governance/semantic_patterns.py @@ -0,0 +1,473 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Discover semantic proof, ownership, lifecycle, and boundary risks in tests.""" + +from __future__ import annotations + +import ast + +from .ast_analysis import call_name, configured_call_name +from .model import TestCandidate + +DRAIN_RULE = "DRAIN001" +OPTIONAL_PROOF_RULE = "OPTIONAL001" +UNBOUNDED_WAIT_RULE = "BOUND001" +SUPPRESSED_FAILURE_RULE = "SUPPRESS001" +RANDOMNESS_RULE = "RANDOM001" +NETWORK_RULE = "NETWORK001" + +_QUEUED_TURN_CALLS = frozenset({"wait_for_queued_qt_turn"}) +_OPTIONAL_PROOF_CALLS = frozenset( + { + "pytest.importorskip", + "pytest.skip", + "pytest.xfail", + } +) +_OPTIONAL_PROOF_MARKERS = frozenset( + { + "pytest.mark.dependency", + "pytest.mark.flaky", + "pytest.mark.order", + "pytest.mark.run", + "pytest.mark.skip", + "pytest.mark.skipif", + "pytest.mark.xfail", + } +) +_EXTERNAL_CALLS_REQUIRING_TIMEOUT = frozenset( + { + "requests.delete", + "requests.get", + "requests.patch", + "requests.post", + "requests.put", + "subprocess.call", + "subprocess.check_call", + "subprocess.check_output", + "subprocess.run", + } +) +_WAIT_METHODS_REQUIRING_BOUND = frozenset({"communicate", "join", "wait"}) +_GLOBAL_RANDOM_CALLS = frozenset( + { + "random.choice", + "random.choices", + "random.getrandbits", + "random.randint", + "random.random", + "random.randrange", + "random.sample", + "random.shuffle", + "random.uniform", + } +) +_REAL_NETWORK_CALLS = frozenset( + { + "httpx.delete", + "httpx.get", + "httpx.patch", + "httpx.post", + "httpx.put", + "httpx.request", + "requests.delete", + "requests.get", + "requests.patch", + "requests.post", + "requests.put", + "requests.request", + "urllib.request.urlopen", + } +) + + +def semantic_pattern_candidates( + *, + relative_path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Return every semantic-risk candidate discovered in one test source.""" + + return [ + *_count_shaped_queued_turn_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + *_manual_qt_event_poll_candidates(path=relative_path, tree=tree), + *_optional_proof_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + *_unbounded_external_wait_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + *_suppressed_failure_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + *_uncontrolled_randomness_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + *_real_network_candidates( + path=relative_path, + tree=tree, + aliases=aliases, + ), + ] + + +def _real_network_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find direct network access outside an explicit qualification owner.""" + + if path.startswith("tests/qualification/"): + return [] + candidates: list[TestCandidate] = [] + ordinal = 0 + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + identity = call_name(node.func, aliases) + if identity not in _REAL_NETWORK_CALLS: + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=NETWORK_RULE, + path=path, + locator=f":real-network:{ordinal}", + evidence=f"calls {identity} outside a qualification owner", + line=node.lineno, + ) + ) + return candidates + + +def _suppressed_failure_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find broad exceptions converted directly into silent control flow.""" + + candidates: list[TestCandidate] = [] + ordinal = 0 + for node in ast.walk(tree): + broad_suppress = ( + isinstance(node, ast.Call) + and call_name(node.func, aliases) == "contextlib.suppress" + and any(_exception_name(argument) for argument in node.args) + ) + silent_handler = ( + isinstance(node, ast.ExceptHandler) + and _exception_name(node.type) + and all( + isinstance(statement, (ast.Pass, ast.Return, ast.Continue)) + for statement in node.body + ) + ) + if not broad_suppress and not silent_handler: + continue + ordinal += 1 + line = node.lineno if isinstance(node, (ast.Call, ast.ExceptHandler)) else 1 + candidates.append( + TestCandidate( + rule=SUPPRESSED_FAILURE_RULE, + path=path, + locator=f":suppressed-failure:{ordinal}", + evidence="broad exception can complete without observable failure", + line=line, + ) + ) + return candidates + + +def _exception_name(node: ast.expr | None) -> bool: + """Return whether one expression names Exception or BaseException.""" + + if isinstance(node, ast.Name): + return node.id in {"Exception", "BaseException"} + if isinstance(node, ast.Tuple): + return any(_exception_name(element) for element in node.elts) + return False + + +def _uncontrolled_randomness_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find random behavior whose seed cannot be replayed from test evidence.""" + + candidates: list[TestCandidate] = [] + ordinal = 0 + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + identity = call_name(node.func, aliases) + unseeded_instance = ( + identity == "random.Random" and not node.args and not node.keywords + ) + if identity not in _GLOBAL_RANDOM_CALLS and not unseeded_instance: + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=RANDOMNESS_RULE, + path=path, + locator=f":uncontrolled-randomness:{ordinal}", + evidence=f"calls {identity} without a replayable local seed", + line=node.lineno, + ) + ) + return candidates + + +def _unbounded_external_wait_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find external operations whose failure path has no explicit time bound.""" + + candidates: list[TestCandidate] = [] + ordinal = 0 + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + identity = call_name(node.func, aliases) + external_call = identity in _EXTERNAL_CALLS_REQUIRING_TIMEOUT + wait_method = ( + isinstance(node.func, ast.Attribute) + and node.func.attr in _WAIT_METHODS_REQUIRING_BOUND + ) + if not external_call and not wait_method: + continue + has_timeout_keyword = any(keyword.arg == "timeout" for keyword in node.keywords) + if external_call and has_timeout_keyword: + continue + if wait_method and (node.args or has_timeout_keyword): + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=UNBOUNDED_WAIT_RULE, + path=path, + locator=f":unbounded-external-wait:{ordinal}", + evidence=f"calls {identity} without an explicit failure bound", + line=node.lineno, + ) + ) + return candidates + + +def _optional_proof_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find skips, expected failures, retries, and order-dependent proof.""" + + parent_by_node = { + child: parent + for parent in ast.walk(tree) + for child in ast.iter_child_nodes(parent) + } + candidates: list[TestCandidate] = [] + ordinal = 0 + for node in ast.walk(tree): + if isinstance(node, ast.Call): + identity = call_name(node.func, aliases) + if identity not in _OPTIONAL_PROOF_CALLS | _OPTIONAL_PROOF_MARKERS: + continue + name = identity + elif isinstance(node, ast.Attribute): + name = call_name(node, aliases) + if name not in _OPTIONAL_PROOF_MARKERS: + continue + parent = parent_by_node.get(node) + if isinstance(parent, ast.Call) and parent.func is node: + continue + else: + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=OPTIONAL_PROOF_RULE, + path=path, + locator=f":optional-proof:{ordinal}", + evidence=f"uses {name}, which can suppress, retry, or reorder proof", + line=node.lineno, + ) + ) + return candidates + + +def _count_shaped_queued_turn_candidates( + *, + path: str, + tree: ast.Module, + aliases: dict[str, str], +) -> list[TestCandidate]: + """Find numeric settling parameters that cannot control queued-turn delivery.""" + + candidates: list[TestCandidate] = [] + for function in ( + node + for node in ast.walk(tree) + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + ): + if not any( + isinstance(node, ast.Call) + and configured_call_name( + call_name(node.func, aliases), + _QUEUED_TURN_CALLS, + ) + is not None + for node in ast.walk(function) + ): + continue + parent_by_node = { + child: parent + for parent in ast.walk(function) + for child in ast.iter_child_nodes(parent) + } + for parameter in _integer_parameters(function): + references = [ + node + for node in ast.walk(function) + if isinstance(node, ast.Name) + and isinstance(node.ctx, ast.Load) + and node.id == parameter.arg + ] + if any( + not _is_ignored_or_boolean_guard(reference, parent_by_node) + for reference in references + ): + continue + candidates.append( + TestCandidate( + rule=DRAIN_RULE, + path=path, + locator=f"{function.name}:count-shaped-queued-turn:{parameter.arg}", + evidence=( + "numeric settling parameter does not control repetition " + "around one queued-turn barrier" + ), + line=function.lineno, + ) + ) + return candidates + + +def _manual_qt_event_poll_candidates( + *, + path: str, + tree: ast.Module, +) -> list[TestCandidate]: + """Find polling loops that manually pump Qt instead of observing owner state.""" + + candidates: list[TestCandidate] = [] + ordinal = 0 + for loop in (node for node in ast.walk(tree) if isinstance(node, ast.While)): + process_events_call = next( + ( + node + for node in ast.walk(loop) + if isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr == "processEvents" + ), + None, + ) + if process_events_call is None: + continue + ordinal += 1 + candidates.append( + TestCandidate( + rule=DRAIN_RULE, + path=path, + locator=f":manual-qt-event-poll:{ordinal}", + evidence=( + "polling loop manually pumps processEvents instead of waiting " + "on a bounded observable owner condition" + ), + line=loop.lineno, + ) + ) + return candidates + + +def _integer_parameters( + function: ast.FunctionDef | ast.AsyncFunctionDef, +) -> tuple[ast.arg, ...]: + """Return parameters whose annotation declares an integer contract.""" + + parameters = ( + *function.args.posonlyargs, + *function.args.args, + *function.args.kwonlyargs, + ) + return tuple( + parameter + for parameter in parameters + if parameter.annotation is not None + and any( + isinstance(node, ast.Name) and node.id == "int" + for node in ast.walk(parameter.annotation) + ) + ) + + +def _is_ignored_or_boolean_guard( + reference: ast.Name, + parent_by_node: dict[ast.AST, ast.AST], +) -> bool: + """Return whether one parameter reference only discards or gates one barrier.""" + + node: ast.AST = reference + while node in parent_by_node: + parent = parent_by_node[node] + if isinstance(parent, ast.Assign): + return any( + isinstance(target, ast.Name) and target.id == "_" + for target in parent.targets + ) + if isinstance(parent, (ast.If, ast.IfExp)): + return node is parent.test + if isinstance( + parent, (ast.BoolOp, ast.Compare, ast.UnaryOp, ast.Tuple, ast.List) + ): + node = parent + continue + return False + return False + + +__all__ = [ + "DRAIN_RULE", + "NETWORK_RULE", + "OPTIONAL_PROOF_RULE", + "RANDOMNESS_RULE", + "SUPPRESSED_FAILURE_RULE", + "UNBOUNDED_WAIT_RULE", + "semantic_pattern_candidates", +] diff --git a/tools/test_governance/validation.py b/tools/test_governance/validation.py new file mode 100644 index 0000000..6786219 --- /dev/null +++ b/tools/test_governance/validation.py @@ -0,0 +1,404 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Validate discovered test candidates against exact reviewed state.""" + +from __future__ import annotations + +from collections import Counter +from datetime import UTC, date, datetime +from pathlib import Path + +from tools.architecture_governance.model import Diagnostic + +from .discovery import discover_test_candidates +from .loading import load_test_policy, load_test_state +from .metrics import reviewed_state_fingerprint +from .model import ( + TestCandidate, + TestDebt, + TestPolicy, + TestState, + TestValidationResult, + TestWaiver, +) + +_MINIMUM_RATIONALE_LENGTH = 120 +_CLASSIFICATION_DISPOSITIONS = frozenset( + { + "false_positive", + "framework_infrastructure", + "intentional_real_time", + "performance_qualification", + "platform_native", + "process_isolated", + "resource_locked", + } +) + + +def validate_test_governance( + root: Path, + *, + policy_path: Path | None = None, + today: date | None = None, +) -> TestValidationResult: + """Return discovered candidates and every current governance diagnostic.""" + + try: + policy = load_test_policy( + policy_path or root / "governance/testing/policy.toml" + ) + state = load_test_state(root, policy) + candidates = discover_test_candidates(root, policy) + except (OSError, SyntaxError, TypeError, ValueError) as error: + return TestValidationResult( + candidates=(), + diagnostics=( + Diagnostic("TSTATE001", "governance/testing/policy.toml", str(error)), + ), + ) + current_date = today or datetime.now(UTC).date() + diagnostics = [ + *_validate_policy(root, policy), + *_validate_unique_state(state), + *_validate_state(root, policy, state, candidates, current_date), + ] + return TestValidationResult( + candidates=candidates, + diagnostics=tuple( + sorted( + diagnostics, + key=lambda item: (item.path, item.rule, item.severity, item.message), + ) + ), + ) + + +def _validate_policy(root: Path, policy: TestPolicy) -> list[Diagnostic]: + """Validate exact policy paths and non-overlapping source declarations.""" + + diagnostics: list[Diagnostic] = [] + if not (root / policy.test_root).is_dir(): + diagnostics.append( + Diagnostic( + "TPOLICY001", "governance/testing/policy.toml", "test_root must exist" + ) + ) + for support_root in policy.semantic_support_roots: + if not (root / support_root).is_dir(): + diagnostics.append( + Diagnostic( + "TPOLICY004", + "governance/testing/policy.toml", + f"semantic support root {support_root.as_posix()} must exist", + ) + ) + if not (root / policy.serial_policy).is_file(): + diagnostics.append( + Diagnostic( + "TPOLICY002", + "governance/testing/policy.toml", + "serial_policy must exist", + ) + ) + for allowed_path in sorted(policy.allowed_root_source_paths): + path = root / allowed_path + if not path.is_file() or path.parent != root / policy.test_root: + diagnostics.append( + Diagnostic( + "TPOLICY003", + "governance/testing/policy.toml", + f"allowed root source {allowed_path} must exist directly " + "under the test root", + ) + ) + return diagnostics + + +def _validate_unique_state(state: TestState) -> list[Diagnostic]: + """Reject duplicate identifiers and candidate dispositions.""" + + diagnostics: list[Diagnostic] = [] + identifiers = [ + *(debt.identifier for debt in state.debts), + *(waiver.identifier for waiver in state.waivers), + ] + for identifier, count in Counter(identifiers).items(): + if count > 1: + diagnostics.append( + Diagnostic( + "TSTATE002", + "governance/testing/policy.toml", + f"test-governance record id {identifier} is not unique", + ) + ) + classified = [ + candidate for waiver in state.waivers for candidate in waiver.candidates + ] + for candidate, count in Counter(classified).items(): + if count > 1: + diagnostics.append( + Diagnostic( + "TSTATE003", + "governance/testing/waivers.toml", + f"candidate {candidate} has multiple dispositions", + ) + ) + debt_candidates = [ + candidate for debt in state.debts for candidate in debt.candidates + ] + for candidate, count in Counter(debt_candidates).items(): + if count > 1: + diagnostics.append( + Diagnostic( + "TSTATE004", + "governance/testing/debt.toml", + f"candidate {candidate} appears in multiple debt records", + ) + ) + return diagnostics + + +def _validate_state( + root: Path, + policy: TestPolicy, + state: TestState, + candidates: tuple[TestCandidate, ...], + today: date, +) -> list[Diagnostic]: + """Validate reviewed records and require one disposition per candidate.""" + + diagnostics: list[Diagnostic] = [] + candidates_by_key = {candidate.key: candidate for candidate in candidates} + debt_by_id = {debt.identifier: debt for debt in state.debts} + for debt in state.debts: + diagnostics.extend( + _validate_debt(root, policy, debt, candidates_by_key, state, today) + ) + for waiver in state.waivers: + diagnostics.extend( + _validate_waiver( + root, + policy, + waiver, + candidates_by_key, + debt_by_id, + today, + ) + ) + classified = { + candidate for waiver in state.waivers for candidate in waiver.candidates + } + for candidate in candidates: + if candidate.key in classified: + continue + diagnostics.append( + Diagnostic( + candidate.rule, + candidate.path, + f"{candidate.evidence} at {candidate.locator}; perform " + "source-level review and record an exact classification or " + "debt-remediation waiver", + ) + ) + return diagnostics + + +def _validate_debt( + root: Path, + policy: TestPolicy, + debt: TestDebt, + candidates: dict[str, TestCandidate], + state: TestState, + today: date, +) -> list[Diagnostic]: + """Validate one exact current test-debt assessment.""" + + registry = policy.debt_registry.as_posix() + diagnostics = _validate_record_identity( + root, + registry, + debt.rule, + debt.candidates, + debt.paths, + debt.fingerprint, + candidates, + "TDEBT", + ) + if debt.review_by < today: + diagnostics.append( + Diagnostic( + "TDEBT004", + registry, + f"debt {debt.identifier} expired on {debt.review_by.isoformat()}", + ) + ) + linked = tuple( + waiver + for waiver in state.waivers + if waiver.kind == "remediation" and waiver.debt == debt.identifier + ) + if len(linked) != 1: + diagnostics.append( + Diagnostic( + "TDEBT005", + registry, + f"debt {debt.identifier} must have exactly one remediation waiver; " + f"found {len(linked)}", + ) + ) + elif ( + linked[0].rule != debt.rule + or linked[0].candidates != debt.candidates + or linked[0].paths != debt.paths + or linked[0].fingerprint != debt.fingerprint + ): + diagnostics.append( + Diagnostic( + "TDEBT006", + registry, + f"debt {debt.identifier} and its remediation waiver must cover " + "identical state", + ) + ) + return diagnostics + + +def _validate_waiver( + root: Path, + policy: TestPolicy, + waiver: TestWaiver, + candidates: dict[str, TestCandidate], + debts: dict[str, TestDebt], + today: date, +) -> list[Diagnostic]: + """Validate one reviewed classification or remediation waiver.""" + + registry = policy.waiver_registry.as_posix() + diagnostics = _validate_record_identity( + root, + registry, + waiver.rule, + waiver.candidates, + waiver.paths, + waiver.fingerprint, + candidates, + "TWAIVER", + ) + if waiver.review_by < today: + diagnostics.append( + Diagnostic( + "TWAIVER004", + registry, + f"waiver {waiver.identifier} expired on {waiver.review_by.isoformat()}", + ) + ) + if len(waiver.rationale) < _MINIMUM_RATIONALE_LENGTH: + diagnostics.append( + Diagnostic( + "TWAIVER005", + registry, + f"waiver {waiver.identifier} requires a substantive " + "source-specific rationale", + ) + ) + if waiver.kind == "classification": + if waiver.disposition not in _CLASSIFICATION_DISPOSITIONS: + diagnostics.append( + Diagnostic( + "TWAIVER006", + registry, + f"waiver {waiver.identifier} uses unsupported classification " + f"{waiver.disposition}", + ) + ) + else: + if waiver.disposition != "debt": + diagnostics.append( + Diagnostic( + "TWAIVER007", + registry, + f"remediation waiver {waiver.identifier} must use disposition debt", + ) + ) + if waiver.debt not in debts: + diagnostics.append( + Diagnostic( + "TWAIVER008", + registry, + f"waiver {waiver.identifier} must link existing test debt", + ) + ) + return diagnostics + + +def _validate_record_identity( + root: Path, + registry: str, + rule: str, + candidate_keys: tuple[str, ...], + paths: tuple[str, ...], + fingerprint: str, + candidates: dict[str, TestCandidate], + prefix: str, +) -> list[Diagnostic]: + """Validate candidate, path, ordering, and source identity for one record.""" + + diagnostics: list[Diagnostic] = [] + exact_candidates = tuple( + candidates[key] for key in candidate_keys if key in candidates + ) + if len(exact_candidates) != len(candidate_keys): + diagnostics.append( + Diagnostic( + f"{prefix}001", + registry, + "record must reference exact current candidate keys", + ) + ) + return diagnostics + if any(candidate.rule != rule for candidate in exact_candidates): + diagnostics.append( + Diagnostic( + f"{prefix}002", + registry, + f"record rule {rule} does not match all referenced candidates", + ) + ) + expected_paths = tuple(sorted({candidate.path for candidate in exact_candidates})) + if paths != expected_paths: + diagnostics.append( + Diagnostic( + f"{prefix}003", + registry, + "record paths must exactly equal sorted candidate paths " + f"{expected_paths}", + ) + ) + elif ( + reviewed_state_fingerprint( + root, + rule=rule, + candidates=candidate_keys, + paths=paths, + ) + != fingerprint + ): + diagnostics.append( + Diagnostic( + f"{prefix}009", + registry, + "record source fingerprint is stale; repeat human review", + ) + ) + if candidate_keys != tuple(sorted(candidate_keys)): + diagnostics.append( + Diagnostic( + f"{prefix}010", + registry, + "record candidates must use stable sorted order", + ) + ) + return diagnostics diff --git a/web/dist/simple-syrup.js b/web/dist/simple-syrup.js index fbd1b7f..6e64526 100644 --- a/web/dist/simple-syrup.js +++ b/web/dist/simple-syrup.js @@ -1,40 +1,21 @@ // web/src/main.ts import { app } from "../../../scripts/app.js"; +// web/src/apiTransport.ts +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/api.ts var SETTINGS_ROUTE = "/simple-syrup/settings"; -var QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache"; -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"; -var MASK_BATCH_PREVIEW_ROUTE = "/simple-syrup/mask-batch/preview"; -async function getMaskBatchPreview(files, channel, fetchImpl = fetch) { - const response = await fetchImpl(MASK_BATCH_PREVIEW_ROUTE, { - method: "POST", - headers: { "Content-Type": "application/json" }, - body: JSON.stringify({ files, channel }) - }); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not render Load Mask Batch preview. Backend returned ${String(response.status)}.` - ) - ); - } - return parseMaskBatchPreview(await response.json()); -} -function parseMaskBatchPreview(payload) { - if (!isMaskBatchPreviewPayload(payload)) { - throw new Error( - "SimpleSyrup mask batch preview payload is invalid. Expected native images and animation flags." - ); - } - return { - images: payload.images.map((image) => ({ ...image })), - animated: [...payload.animated] - }; -} async function getSettings(fetchImpl = fetch) { const response = await fetchImpl(SETTINGS_ROUTE); if (!response.ok) { @@ -74,107 +55,55 @@ function parseSettings(payload) { quant_cache_limit_gib: payload.quant_cache_limit_gib }; } -async function getQuantCacheStatus(fetchImpl = fetch) { - const response = await fetchImpl(QUANT_CACHE_ROUTE); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not load quant cache status. Backend returned ${String(response.status)}.` - ) - ); - } - return parseQuantCacheStatus(await response.json()); -} -async function clearQuantCache(fetchImpl = fetch) { - const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "DELETE" }); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not clear quant cache. Backend returned ${String(response.status)}.` - ) - ); - } - return parseQuantCacheStatus(await response.json()); -} -async function enforceQuantCacheLimit(fetchImpl = fetch) { - const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "POST" }); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not enforce quant cache limit. Backend returned ${String(response.status)}.` - ) - ); - } - return parseQuantCacheStatus(await response.json()); -} -function parseQuantCacheStatus(payload) { - if (!isQuantCacheStatusPayload(payload)) { - throw new Error( - "SimpleSyrup quant cache status is invalid. Expected path, byte usage, limit, and artifact counts." - ); - } - return { ...payload }; +function isSettingsPayload(payload) { + return typeof payload === "object" && payload !== null && typeof payload.show_downloadable_models === "boolean" && Number.isInteger( + payload.quant_cache_limit_gib + ) && Number(payload.quant_cache_limit_gib) > 0; } + +// web/src/externalLlmApi.ts +var SETTINGS_ROUTE2 = "/simple-syrup/external-llm/settings"; +var API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"; +var MODELS_REFRESH_ROUTE = "/simple-syrup/external-llm/models/refresh"; 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()); + return requestSettings( + fetchImpl, + SETTINGS_ROUTE2, + void 0, + "Could not load external LLM settings" + ); } 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()); + return requestSettings( + fetchImpl, + SETTINGS_ROUTE2, + { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(settings) + }, + "Could not save external LLM settings" + ); } 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()); + return requestSettings( + fetchImpl, + API_KEY_ROUTE, + { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(payload) + }, + "Could not save external LLM API key" + ); } 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()); + return requestSettings( + fetchImpl, + MODELS_REFRESH_ROUTE, + { method: "POST" }, + "Could not refresh external LLM models" + ); } function parseExternalLLMSettings(payload) { if (!isExternalLLMSettingsPayload(payload)) { @@ -189,10 +118,52 @@ function parseExternalLLMSettings(payload) { has_api_key: payload.has_api_key }; } -function isSettingsPayload(payload) { - return typeof payload === "object" && payload !== null && typeof payload.show_downloadable_models === "boolean" && Number.isInteger( - payload.quant_cache_limit_gib - ) && Number(payload.quant_cache_limit_gib) > 0; +async function requestSettings(fetchImpl, route, init, errorPrefix) { + const response = init === void 0 ? await fetchImpl(route) : await fetchImpl(route, init); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `${errorPrefix}. Backend returned ${String(response.status)}.` + ) + ); + } + return parseExternalLLMSettings(await response.json()); +} +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"; +} + +// web/src/quantCacheApi.ts +var QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache"; +async function getQuantCacheStatus(fetchImpl = fetch) { + return requestQuantCache(fetchImpl); +} +async function clearQuantCache(fetchImpl = fetch) { + return requestQuantCache(fetchImpl, "DELETE"); +} +async function enforceQuantCacheLimit(fetchImpl = fetch) { + return requestQuantCache(fetchImpl, "POST"); +} +function parseQuantCacheStatus(payload) { + if (!isQuantCacheStatusPayload(payload)) { + throw new Error( + "SimpleSyrup quant cache status is invalid. Expected path, byte usage, limit, and artifact counts." + ); + } + return { ...payload }; +} +async function requestQuantCache(fetchImpl, method) { + const response = method === void 0 ? await fetchImpl(QUANT_CACHE_ROUTE) : await fetchImpl(QUANT_CACHE_ROUTE, { method }); + if (!response.ok) { + const fallback = method === "DELETE" ? `Could not clear quant cache. Backend returned ${String(response.status)}.` : method === "POST" ? `Could not enforce quant cache limit. Backend returned ${String(response.status)}.` : `Could not load quant cache status. Backend returned ${String(response.status)}.`; + throw new Error( + await backendErrorMessage(response, fallback) + ); + } + return parseQuantCacheStatus(await response.json()); } function isQuantCacheStatusPayload(payload) { if (typeof payload !== "object" || payload === null) return false; @@ -202,32 +173,6 @@ function isQuantCacheStatusPayload(payload) { function isNonNegativeInteger(value) { return typeof value === "number" && Number.isInteger(value) && value >= 0; } -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"; -} -function isMaskBatchPreviewPayload(payload) { - if (typeof payload !== "object" || payload === null) return false; - const candidate = payload; - return Array.isArray(candidate.images) && candidate.images.every(isComfyImageResult) && Array.isArray(candidate.animated) && candidate.animated.every((value) => typeof value === "boolean"); -} -function isComfyImageResult(value) { - if (typeof value !== "object" || value === null) return false; - const candidate = value; - return typeof candidate.filename === "string" && typeof candidate.subfolder === "string" && (candidate.type === "input" || candidate.type === "output" || candidate.type === "temp"); -} -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/downloadableModelsSetting.ts var SIMPLE_SYRUP_SETTING_ID = "SimpleSyrup.ShowDownloadableModels"; @@ -745,6 +690,46 @@ function registerExternalLLMRefreshHook(app2, api = { refreshExternalLLMModels } }; } +// web/src/maskBatchPreviewApi.ts +var MASK_BATCH_PREVIEW_ROUTE = "/simple-syrup/mask-batch/preview"; +async function getMaskBatchPreview(files, channel, fetchImpl = fetch) { + const response = await fetchImpl(MASK_BATCH_PREVIEW_ROUTE, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ files, channel }) + }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not render Load Mask Batch preview. Backend returned ${String(response.status)}.` + ) + ); + } + return parseMaskBatchPreview(await response.json()); +} +function parseMaskBatchPreview(payload) { + if (!isMaskBatchPreviewPayload(payload)) { + throw new Error( + "SimpleSyrup mask batch preview payload is invalid. Expected native images and animation flags." + ); + } + return { + images: payload.images.map((image) => ({ ...image })), + animated: [...payload.animated] + }; +} +function isMaskBatchPreviewPayload(payload) { + if (typeof payload !== "object" || payload === null) return false; + const candidate = payload; + return Array.isArray(candidate.images) && candidate.images.every(isComfyImageResult) && Array.isArray(candidate.animated) && candidate.animated.every((value) => typeof value === "boolean"); +} +function isComfyImageResult(value) { + if (typeof value !== "object" || value === null) return false; + const candidate = value; + return typeof candidate.filename === "string" && typeof candidate.subfolder === "string" && (candidate.type === "input" || candidate.type === "output" || candidate.type === "temp"); +} + // web/src/nativeNodePreview.ts var NativeNodePreview = class { constructor(app2, api, node) { @@ -1235,6 +1220,110 @@ var OrderedMediaPreviewTransaction = class { } }; +// web/src/orderedMediaDetailSelection.ts +var OrderedMediaDetailSelection = class { + constructor(options) { + this.options = options; + } + options; + pendingIndex = null; + restoreFrame = null; + /** Keep the moved item selected across both native renderer state models. */ + followMoved(index, destination) { + if (this.options.selectedIndex() !== index) return; + this.pendingIndex = destination; + if (this.options.canvasIndex() === index) { + this.options.setCanvasIndex(destination); + } else { + this.options.selectDomIndex(destination); + } + this.options.refreshActions(); + } + /** Select the nearest remaining item after removing an inspected item. */ + followRemoved(index, removedDetail) { + if (!removedDetail) return; + const remaining = this.options.itemCount(); + const destination = remaining > 0 ? Math.min(index, remaining - 1) : null; + this.pendingIndex = destination; + if (this.options.canvasIndex() === index) { + this.options.setCanvasIndex(destination); + } else if (destination !== null) { + this.options.selectDomIndex(destination); + } + this.options.refreshActions(); + } + /** Re-enter Nodes 2.0 detail mode after output publication resets its grid. */ + restorePending() { + if (this.pendingIndex === null || this.restoreFrame !== null) return; + let stableFrames = 0; + let remainingFrames = 12; + const restore = () => { + this.restoreFrame = null; + const destination = this.pendingIndex; + if (destination === null) return; + if (this.options.domSelectedIndex() === destination) { + stableFrames += 1; + } else { + stableFrames = 0; + this.options.selectDomIndex(destination); + } + remainingFrames -= 1; + if (stableFrames >= 3 || remainingFrames <= 0) { + this.pendingIndex = null; + this.options.refreshActions(); + return; + } + this.restoreFrame = requestAnimationFrame(restore); + }; + this.restoreFrame = requestAnimationFrame(restore); + } + /** Cancel any selection restoration still waiting for native publication. */ + dispose() { + if (this.restoreFrame === null) return; + cancelAnimationFrame(this.restoreFrame); + this.restoreFrame = null; + } +}; + +// web/src/orderedMediaPreviewGeometry.ts +function previewItems(images) { + return images.map((image) => ({ + sourceUrl: image.currentSrc || image.src, + image + })); +} +function imageSlot(image) { + return () => elementSlot(image); +} +function elementSlot(element) { + const rect = element.getBoundingClientRect(); + return { + left: rect.left, + top: rect.top, + width: rect.width, + height: rect.height + }; +} +function validSlot(slot) { + return Number.isFinite(slot.left) && Number.isFinite(slot.top) && Number.isFinite(slot.width) && Number.isFinite(slot.height) && slot.width >= 0 && slot.height >= 0; +} +function imageArea(image) { + const rect = image.getBoundingClientRect(); + return rect.width * rect.height; +} +function indexedSlots(slots, container = null) { + return slots.map( + (bounds, itemIndex) => container ? { itemIndex, bounds, container } : { itemIndex, bounds } + ); +} +function unionImageRects(rects) { + const left = Math.min(...rects.map(([x]) => x)); + const top = Math.min(...rects.map(([, y]) => y)); + const right = Math.max(...rects.map(([x, , width]) => x + width)); + const bottom = Math.max(...rects.map(([, y, , height]) => y + height)); + return [left, top, right - left, bottom - top]; +} + // web/src/comfyImageReference.ts function comfyImageReferenceKey(reference) { return `${reference.type} @@ -1275,20 +1364,20 @@ var OrderedMediaPreviewActions = class { const moveEarlier = (index) => { const destination = index - 1; this.transaction.move(index, destination); - this.followMovedDetail(index, destination); + this.detailSelection.followMoved(index, destination); options.moveEarlier(index); }; const moveLater = (index) => { const destination = index + 1; this.transaction.move(index, destination); - this.followMovedDetail(index, destination); + this.detailSelection.followMoved(index, destination); options.moveLater(index); }; const remove = (index) => { const removedDetail = this.selectedItemIndex() === index; this.transaction.remove(index); options.remove(index); - this.followRemovedDetail(index, removedDetail); + this.detailSelection.followRemoved(index, removedDetail); }; const actionOptions = { getItemCount: () => this.itemCount(), @@ -1301,10 +1390,30 @@ var OrderedMediaPreviewActions = class { getSlots: () => this.nativeActionSlots(), ...actionOptions }); + this.detailSelection = new OrderedMediaDetailSelection({ + selectedIndex: () => this.selectedItemIndex(), + canvasIndex: () => this.options.node.imageIndex, + setCanvasIndex: (index) => { + this.options.node.imageIndex = index; + }, + itemCount: () => this.itemCount(), + domSelectedIndex: () => { + const selectedIndex = this.domDetailButtons().findIndex( + (button) => button.getAttribute("aria-current") === "true" + ); + return selectedIndex >= 0 ? selectedIndex : null; + }, + selectDomIndex: (index) => { + this.domDetailButtons()[index]?.click(); + }, + refreshActions: () => { + this.affordances.refresh(); + } + }); this.unsubscribePreview = options.preview.subscribe(() => { this.affordances.refresh(); this.transaction.authoritativePublished(); - this.restorePendingDetail(); + this.detailSelection.restorePending(); }); this.unsubscribeLifecycle = subscribeNativePreviewLifecycle(() => { this.affordances.refresh(); @@ -1313,21 +1422,17 @@ var OrderedMediaPreviewActions = class { options; affordances; transaction; + detailSelection; unsubscribePreview; unsubscribeLifecycle; lastCanvasPreviewRect = null; - pendingDetailIndex = null; - detailRestoreFrame = null; /** Remove layout listeners and every loader-owned overlay control. */ dispose() { this.unsubscribePreview(); this.unsubscribeLifecycle(); this.transaction.dispose(); this.affordances.dispose(); - if (this.detailRestoreFrame !== null) { - cancelAnimationFrame(this.detailRestoreFrame); - this.detailRestoreFrame = null; - } + this.detailSelection.dispose(); } itemCount() { const files = this.options.getFiles(); @@ -1621,61 +1726,6 @@ var OrderedMediaPreviewActions = class { const detailIndex = this.domDetailButtons().indexOf(currentButton); return detailIndex >= 0 && detailIndex < this.itemCount() ? detailIndex : null; } - /** Keep the moved item selected across both native renderer state models. */ - followMovedDetail(index, destination) { - if (this.selectedItemIndex() !== index) return; - this.pendingDetailIndex = destination; - if (this.options.node.imageIndex === index) { - this.options.node.imageIndex = destination; - } else { - this.domDetailButtons()[destination]?.click(); - } - this.affordances.refresh(); - } - /** Select the nearest remaining item after removing an inspected item. */ - followRemovedDetail(index, removedDetail) { - if (!removedDetail) return; - const remaining = this.itemCount(); - const destination = remaining > 0 ? Math.min(index, remaining - 1) : null; - this.pendingDetailIndex = destination; - if (this.options.node.imageIndex === index) { - this.options.node.imageIndex = destination; - } else if (destination !== null) { - this.domDetailButtons()[destination]?.click(); - } - this.affordances.refresh(); - } - /** Re-enter Nodes 2.0 detail mode after output publication resets its grid. */ - restorePendingDetail() { - if (this.pendingDetailIndex === null || this.detailRestoreFrame !== null) { - return; - } - let stableFrames = 0; - let remainingFrames = 12; - const restore = () => { - this.detailRestoreFrame = null; - const destination = this.pendingDetailIndex; - if (destination === null) return; - const buttons = this.domDetailButtons(); - const selected = buttons.findIndex( - (button) => button.getAttribute("aria-current") === "true" - ); - if (selected === destination) { - stableFrames += 1; - } else { - stableFrames = 0; - buttons[destination]?.click(); - } - remainingFrames -= 1; - if (stableFrames >= 3 || remainingFrames <= 0) { - this.pendingDetailIndex = null; - this.affordances.refresh(); - return; - } - this.detailRestoreFrame = requestAnimationFrame(restore); - }; - this.detailRestoreFrame = requestAnimationFrame(restore); - } /** Return Comfy's ordered detail navigation controls for this node. */ domDetailButtons() { const root = this.domRoot(); @@ -1691,43 +1741,6 @@ var OrderedMediaPreviewActions = class { return nodeId === void 0 ? void 0 : String(nodeId); } }; -function previewItems(images) { - return images.map((image) => ({ - sourceUrl: image.currentSrc || image.src, - image - })); -} -function imageSlot(image) { - return () => elementSlot(image); -} -function elementSlot(element) { - const rect = element.getBoundingClientRect(); - return { - left: rect.left, - top: rect.top, - width: rect.width, - height: rect.height - }; -} -function validSlot(slot) { - return Number.isFinite(slot.left) && Number.isFinite(slot.top) && Number.isFinite(slot.width) && Number.isFinite(slot.height) && slot.width >= 0 && slot.height >= 0; -} -function imageArea(image) { - const rect = image.getBoundingClientRect(); - return rect.width * rect.height; -} -function indexedSlots(slots, container = null) { - return slots.map( - (bounds, itemIndex) => container ? { itemIndex, bounds, container } : { itemIndex, bounds } - ); -} -function unionImageRects(rects) { - const left = Math.min(...rects.map(([x]) => x)); - const top = Math.min(...rects.map(([, y]) => y)); - const right = Math.max(...rects.map(([x, , width]) => x + width)); - const bottom = Math.max(...rects.map(([, y, , height]) => y + height)); - return [left, top, right - left, bottom - top]; -} // web/src/orderedMediaSelection.ts var OrderedMediaSelection = class { diff --git a/web/src/api.ts b/web/src/api.ts index 9275199..492be40 100644 --- a/web/src/api.ts +++ b/web/src/api.ts @@ -2,86 +2,16 @@ // Copyright (C) 2026 Artificial Sweetener and contributors // SPDX-License-Identifier: AGPL-3.0-or-later -import type { ComfyImageResult, ComfyNodeExecutionOutput } from "./types"; +/** General SimpleSyrup settings transport and payload validation. */ + +import { backendErrorMessage, type FetchLike } from "./apiTransport"; export interface SimpleSyrupSettings { show_downloadable_models: boolean; quant_cache_limit_gib: number; } -export interface QuantCacheStatus { - path: string; - usage_bytes: number; - limit_bytes: number; - artifact_count: number; - active_artifact_count: number; - removed_artifacts?: number; - removed_bytes?: number; -} - -export interface ExternalLLMSettings { - base_url: string; - cached_models: string[]; - default_model: string; - has_api_key: boolean; -} - -export interface ExternalLLMSettingsUpdate { - base_url: string; - default_model: string; -} - -export interface ExternalLLMApiKeyUpdate { - api_key: string; -} - -export type FetchLike = ( - input: RequestInfo | URL, - init?: RequestInit -) => Promise; - const SETTINGS_ROUTE = "/simple-syrup/settings"; -const QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache"; -const EXTERNAL_LLM_SETTINGS_ROUTE = "/simple-syrup/external-llm/settings"; -const EXTERNAL_LLM_API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"; -const EXTERNAL_LLM_MODELS_REFRESH_ROUTE = - "/simple-syrup/external-llm/models/refresh"; -const MASK_BATCH_PREVIEW_ROUTE = "/simple-syrup/mask-batch/preview"; - -export async function getMaskBatchPreview( - files: string[], - channel: string, - fetchImpl: FetchLike = fetch -): Promise { - const response = await fetchImpl(MASK_BATCH_PREVIEW_ROUTE, { - method: "POST", - headers: { "Content-Type": "application/json" }, - body: JSON.stringify({ files, channel }) - }); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not render Load Mask Batch preview. Backend returned ${String(response.status)}.` - ) - ); - } - return parseMaskBatchPreview(await response.json()); -} - -export function parseMaskBatchPreview( - payload: unknown -): ComfyNodeExecutionOutput { - if (!isMaskBatchPreviewPayload(payload)) { - throw new Error( - "SimpleSyrup mask batch preview payload is invalid. Expected native images and animation flags." - ); - } - return { - images: payload.images.map((image) => ({ ...image })), - animated: [...payload.animated] - }; -} export async function getSettings( fetchImpl: FetchLike = fetch @@ -130,165 +60,6 @@ export function parseSettings(payload: unknown): SimpleSyrupSettings { }; } -export async function getQuantCacheStatus( - fetchImpl: FetchLike = fetch -): Promise { - const response = await fetchImpl(QUANT_CACHE_ROUTE); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not load quant cache status. Backend returned ${String(response.status)}.` - ) - ); - } - return parseQuantCacheStatus(await response.json()); -} - -export async function clearQuantCache( - fetchImpl: FetchLike = fetch -): Promise { - const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "DELETE" }); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not clear quant cache. Backend returned ${String(response.status)}.` - ) - ); - } - return parseQuantCacheStatus(await response.json()); -} - -export async function enforceQuantCacheLimit( - fetchImpl: FetchLike = fetch -): Promise { - const response = await fetchImpl(QUANT_CACHE_ROUTE, { method: "POST" }); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not enforce quant cache limit. Backend returned ${String(response.status)}.` - ) - ); - } - return parseQuantCacheStatus(await response.json()); -} - -export function parseQuantCacheStatus(payload: unknown): QuantCacheStatus { - if (!isQuantCacheStatusPayload(payload)) { - throw new Error( - "SimpleSyrup quant cache status is invalid. Expected path, byte usage, limit, and artifact counts." - ); - } - return { ...payload }; -} - -export async function getExternalLLMSettings( - fetchImpl: FetchLike = fetch -): Promise { - 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()); -} - -export async function saveExternalLLMSettings( - settings: ExternalLLMSettingsUpdate, - fetchImpl: FetchLike = fetch -): Promise { - 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()); -} - -export async function saveExternalLLMApiKey( - payload: ExternalLLMApiKeyUpdate, - fetchImpl: FetchLike = fetch -): Promise { - 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()); -} - -export async function deleteExternalLLMApiKey( - fetchImpl: FetchLike = fetch -): Promise { - const response = await fetchImpl(EXTERNAL_LLM_API_KEY_ROUTE, { - method: "DELETE" - }); - if (!response.ok) { - throw new Error( - await backendErrorMessage( - response, - `Could not delete external LLM API key. Backend returned ${String(response.status)}.` - ) - ); - } - return parseExternalLLMSettings(await response.json()); -} - -export async function refreshExternalLLMModels( - fetchImpl: FetchLike = fetch -): Promise { - 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()); -} - -export function parseExternalLLMSettings( - payload: unknown -): ExternalLLMSettings { - 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: unknown): payload is SimpleSyrupSettings { return ( typeof payload === "object" && @@ -301,86 +72,3 @@ function isSettingsPayload(payload: unknown): payload is SimpleSyrupSettings { Number((payload as Partial).quant_cache_limit_gib) > 0 ); } - -function isQuantCacheStatusPayload( - payload: unknown -): payload is QuantCacheStatus { - if (typeof payload !== "object" || payload === null) return false; - const candidate = payload as Partial; - return ( - typeof candidate.path === "string" && - isNonNegativeInteger(candidate.usage_bytes) && - isNonNegativeInteger(candidate.limit_bytes) && - isNonNegativeInteger(candidate.artifact_count) && - isNonNegativeInteger(candidate.active_artifact_count) && - (candidate.removed_artifacts === undefined || - isNonNegativeInteger(candidate.removed_artifacts)) && - (candidate.removed_bytes === undefined || - isNonNegativeInteger(candidate.removed_bytes)) - ); -} - -function isNonNegativeInteger(value: unknown): value is number { - return typeof value === "number" && Number.isInteger(value) && value >= 0; -} - -function isExternalLLMSettingsPayload( - payload: unknown -): payload is ExternalLLMSettings { - return ( - typeof payload === "object" && - payload !== null && - typeof (payload as Partial).base_url === "string" && - Array.isArray((payload as Partial).cached_models) && - (payload as Partial).cached_models?.every( - (model) => typeof model === "string" - ) === true && - typeof (payload as Partial).default_model === - "string" && - typeof (payload as Partial).has_api_key === "boolean" - ); -} - -function isMaskBatchPreviewPayload( - payload: unknown -): payload is Required { - if (typeof payload !== "object" || payload === null) return false; - const candidate = payload as Partial; - return ( - Array.isArray(candidate.images) && - candidate.images.every(isComfyImageResult) && - Array.isArray(candidate.animated) && - candidate.animated.every((value) => typeof value === "boolean") - ); -} - -function isComfyImageResult(value: unknown): value is ComfyImageResult { - if (typeof value !== "object" || value === null) return false; - const candidate = value as Partial; - return ( - typeof candidate.filename === "string" && - typeof candidate.subfolder === "string" && - (candidate.type === "input" || - candidate.type === "output" || - candidate.type === "temp") - ); -} - -async function backendErrorMessage( - response: Response, - fallback: string -): Promise { - try { - const payload: unknown = await response.json(); - if ( - typeof payload === "object" && - payload !== null && - typeof (payload as { error?: unknown }).error === "string" - ) { - return (payload as { error: string }).error; - } - } catch { - return fallback; - } - return fallback; -} diff --git a/web/src/apiTransport.ts b/web/src/apiTransport.ts new file mode 100644 index 0000000..2df9e7e --- /dev/null +++ b/web/src/apiTransport.ts @@ -0,0 +1,29 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +/** Shared transport primitives for SimpleSyrup frontend API clients. */ + +export type FetchLike = ( + input: RequestInfo | URL, + init?: RequestInit +) => Promise; + +export async function backendErrorMessage( + response: Response, + fallback: string +): Promise { + try { + const payload: unknown = await response.json(); + if ( + typeof payload === "object" && + payload !== null && + typeof (payload as { error?: unknown }).error === "string" + ) { + return (payload as { error: string }).error; + } + } catch { + return fallback; + } + return fallback; +} diff --git a/web/src/externalLlmApi.ts b/web/src/externalLlmApi.ts new file mode 100644 index 0000000..8ddbcd6 --- /dev/null +++ b/web/src/externalLlmApi.ts @@ -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 + +/** External LLM settings, credential, and model-refresh transport. */ + +import { backendErrorMessage, type FetchLike } from "./apiTransport"; + +export interface ExternalLLMSettings { + base_url: string; + cached_models: string[]; + default_model: string; + has_api_key: boolean; +} + +export interface ExternalLLMSettingsUpdate { + base_url: string; + default_model: string; +} + +export interface ExternalLLMApiKeyUpdate { + api_key: string; +} + +const SETTINGS_ROUTE = "/simple-syrup/external-llm/settings"; +const API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"; +const MODELS_REFRESH_ROUTE = "/simple-syrup/external-llm/models/refresh"; + +export async function getExternalLLMSettings( + fetchImpl: FetchLike = fetch +): Promise { + return requestSettings( + fetchImpl, + SETTINGS_ROUTE, + undefined, + "Could not load external LLM settings" + ); +} + +export async function saveExternalLLMSettings( + settings: ExternalLLMSettingsUpdate, + fetchImpl: FetchLike = fetch +): Promise { + return requestSettings( + fetchImpl, + SETTINGS_ROUTE, + { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(settings) + }, + "Could not save external LLM settings" + ); +} + +export async function saveExternalLLMApiKey( + payload: ExternalLLMApiKeyUpdate, + fetchImpl: FetchLike = fetch +): Promise { + return requestSettings( + fetchImpl, + API_KEY_ROUTE, + { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(payload) + }, + "Could not save external LLM API key" + ); +} + +export async function deleteExternalLLMApiKey( + fetchImpl: FetchLike = fetch +): Promise { + return requestSettings( + fetchImpl, + API_KEY_ROUTE, + { method: "DELETE" }, + "Could not delete external LLM API key" + ); +} + +export async function refreshExternalLLMModels( + fetchImpl: FetchLike = fetch +): Promise { + return requestSettings( + fetchImpl, + MODELS_REFRESH_ROUTE, + { method: "POST" }, + "Could not refresh external LLM models" + ); +} + +export function parseExternalLLMSettings( + payload: unknown +): ExternalLLMSettings { + 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 + }; +} + +async function requestSettings( + fetchImpl: FetchLike, + route: string, + init: RequestInit | undefined, + errorPrefix: string +): Promise { + const response = + init === undefined ? await fetchImpl(route) : await fetchImpl(route, init); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `${errorPrefix}. Backend returned ${String(response.status)}.` + ) + ); + } + return parseExternalLLMSettings(await response.json()); +} + +function isExternalLLMSettingsPayload( + payload: unknown +): payload is ExternalLLMSettings { + return ( + typeof payload === "object" && + payload !== null && + typeof (payload as Partial).base_url === "string" && + Array.isArray((payload as Partial).cached_models) && + (payload as Partial).cached_models?.every( + (model) => typeof model === "string" + ) === true && + typeof (payload as Partial).default_model === "string" && + typeof (payload as Partial).has_api_key === "boolean" + ); +} diff --git a/web/src/externalLlmSettings.ts b/web/src/externalLlmSettings.ts index 7d2738a..2af168a 100644 --- a/web/src/externalLlmSettings.ts +++ b/web/src/externalLlmSettings.ts @@ -8,8 +8,8 @@ import { getExternalLLMSettings, saveExternalLLMApiKey, saveExternalLLMSettings -} from "./api"; -import type { ExternalLLMSettings } from "./api"; +} from "./externalLlmApi"; +import type { ExternalLLMSettings } from "./externalLlmApi"; import { createElement, installSimpleSyrupSettingsStyle, diff --git a/web/src/maskBatchPreviewApi.ts b/web/src/maskBatchPreviewApi.ts new file mode 100644 index 0000000..cddbb12 --- /dev/null +++ b/web/src/maskBatchPreviewApi.ts @@ -0,0 +1,70 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +/** Mask-batch preview transport and native preview payload validation. */ + +import type { ComfyImageResult, ComfyNodeExecutionOutput } from "./types"; +import { backendErrorMessage, type FetchLike } from "./apiTransport"; + +const MASK_BATCH_PREVIEW_ROUTE = "/simple-syrup/mask-batch/preview"; + +export async function getMaskBatchPreview( + files: string[], + channel: string, + fetchImpl: FetchLike = fetch +): Promise { + const response = await fetchImpl(MASK_BATCH_PREVIEW_ROUTE, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ files, channel }) + }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not render Load Mask Batch preview. Backend returned ${String(response.status)}.` + ) + ); + } + return parseMaskBatchPreview(await response.json()); +} + +export function parseMaskBatchPreview( + payload: unknown +): ComfyNodeExecutionOutput { + if (!isMaskBatchPreviewPayload(payload)) { + throw new Error( + "SimpleSyrup mask batch preview payload is invalid. Expected native images and animation flags." + ); + } + return { + images: payload.images.map((image) => ({ ...image })), + animated: [...payload.animated] + }; +} + +function isMaskBatchPreviewPayload( + payload: unknown +): payload is Required { + if (typeof payload !== "object" || payload === null) return false; + const candidate = payload as Partial; + return ( + Array.isArray(candidate.images) && + candidate.images.every(isComfyImageResult) && + Array.isArray(candidate.animated) && + candidate.animated.every((value) => typeof value === "boolean") + ); +} + +function isComfyImageResult(value: unknown): value is ComfyImageResult { + if (typeof value !== "object" || value === null) return false; + const candidate = value as Partial; + return ( + typeof candidate.filename === "string" && + typeof candidate.subfolder === "string" && + (candidate.type === "input" || + candidate.type === "output" || + candidate.type === "temp") + ); +} diff --git a/web/src/maskBatchUpload.ts b/web/src/maskBatchUpload.ts index 308fab3..b57c77e 100644 --- a/web/src/maskBatchUpload.ts +++ b/web/src/maskBatchUpload.ts @@ -2,7 +2,7 @@ // Copyright (C) 2026 Artificial Sweetener and contributors // SPDX-License-Identifier: AGPL-3.0-or-later -import { getMaskBatchPreview } from "./api"; +import { getMaskBatchPreview } from "./maskBatchPreviewApi"; import { configureOrderedMediaNode, registerOrderedMediaNode, diff --git a/web/src/orderedMediaDetailSelection.ts b/web/src/orderedMediaDetailSelection.ts new file mode 100644 index 0000000..7658ddf --- /dev/null +++ b/web/src/orderedMediaDetailSelection.ts @@ -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 + +/** Preserve native detail selection across ordered-media list mutations. */ + +export interface OrderedMediaDetailSelectionOptions { + readonly selectedIndex: () => number | null; + readonly canvasIndex: () => number | null | undefined; + readonly setCanvasIndex: (index: number | null) => void; + readonly itemCount: () => number; + readonly domSelectedIndex: () => number | null; + readonly selectDomIndex: (index: number) => void; + readonly refreshActions: () => void; +} + +export class OrderedMediaDetailSelection { + private pendingIndex: number | null = null; + private restoreFrame: number | null = null; + + constructor(private readonly options: OrderedMediaDetailSelectionOptions) {} + + /** Keep the moved item selected across both native renderer state models. */ + followMoved(index: number, destination: number): void { + if (this.options.selectedIndex() !== index) return; + this.pendingIndex = destination; + if (this.options.canvasIndex() === index) { + this.options.setCanvasIndex(destination); + } else { + this.options.selectDomIndex(destination); + } + this.options.refreshActions(); + } + + /** Select the nearest remaining item after removing an inspected item. */ + followRemoved(index: number, removedDetail: boolean): void { + if (!removedDetail) return; + const remaining = this.options.itemCount(); + const destination = remaining > 0 ? Math.min(index, remaining - 1) : null; + this.pendingIndex = destination; + if (this.options.canvasIndex() === index) { + this.options.setCanvasIndex(destination); + } else if (destination !== null) { + this.options.selectDomIndex(destination); + } + this.options.refreshActions(); + } + + /** Re-enter Nodes 2.0 detail mode after output publication resets its grid. */ + restorePending(): void { + if (this.pendingIndex === null || this.restoreFrame !== null) return; + let stableFrames = 0; + let remainingFrames = 12; + const restore = (): void => { + this.restoreFrame = null; + const destination = this.pendingIndex; + if (destination === null) return; + if (this.options.domSelectedIndex() === destination) { + stableFrames += 1; + } else { + stableFrames = 0; + this.options.selectDomIndex(destination); + } + remainingFrames -= 1; + if (stableFrames >= 3 || remainingFrames <= 0) { + this.pendingIndex = null; + this.options.refreshActions(); + return; + } + this.restoreFrame = requestAnimationFrame(restore); + }; + this.restoreFrame = requestAnimationFrame(restore); + } + + /** Cancel any selection restoration still waiting for native publication. */ + dispose(): void { + if (this.restoreFrame === null) return; + cancelAnimationFrame(this.restoreFrame); + this.restoreFrame = null; + } +} diff --git a/web/src/orderedMediaPreviewActionTypes.ts b/web/src/orderedMediaPreviewActionTypes.ts new file mode 100644 index 0000000..b8e57f5 --- /dev/null +++ b/web/src/orderedMediaPreviewActionTypes.ts @@ -0,0 +1,47 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +/** Define the host-facing collaborators used by ordered-media preview actions. */ + +import type { NativeNodePreview } from "./nativeNodePreview"; +import type { NativePreviewSlot } from "./orderedMediaPreviewAffordances"; +import type { NativeImageRect } from "./orderedMediaPreviewGeometry"; +import type { ComfyApp } from "./types"; + +export interface CubeFaceProjection { + readonly container: HTMLElement; + readonly projectRect: (rect: NativeImageRect) => NativePreviewSlot; +} + +export interface ActionPreviewNode { + readonly id?: string | number; + readonly pos?: readonly [number, number]; + readonly size?: readonly [number, number]; + readonly flags?: { readonly collapsed?: boolean }; + widgets?: ActionPreviewWidget[]; + imageIndex?: number | null; + imageRects?: NativeImageRect[]; + imgs?: HTMLImageElement[]; + graph?: { + _rootGraph?: object; + setDirtyCanvas?: (foreground: boolean, background: boolean) => void; + }; +} + +interface ActionPreviewWidget { + readonly y?: number; + readonly computedHeight?: number; + readonly options?: { readonly canvasOnly?: boolean }; +} + +export interface OrderedMediaPreviewActionOptions { + readonly app: ComfyApp; + readonly node: ActionPreviewNode; + readonly preview: NativeNodePreview; + readonly itemLabel: string; + readonly getFiles: () => string[]; + readonly moveEarlier: (index: number) => void; + readonly moveLater: (index: number) => void; + readonly remove: (index: number) => void; +} diff --git a/web/src/orderedMediaPreviewActions.ts b/web/src/orderedMediaPreviewActions.ts index 67cf2aa..0e8bc17 100644 --- a/web/src/orderedMediaPreviewActions.ts +++ b/web/src/orderedMediaPreviewActions.ts @@ -2,7 +2,6 @@ // Copyright (C) 2026 Artificial Sweetener and contributors // SPDX-License-Identifier: AGPL-3.0-or-later -import type { NativeNodePreview } from "./nativeNodePreview"; import { subscribeNativePreviewLifecycle } from "./nativePreviewLifecycle"; import { OrderedMediaPreviewAffordances, @@ -11,66 +10,43 @@ import { } from "./orderedMediaPreviewAffordances"; import { OrderedMediaPreviewTransaction, - type OrderedMediaPreviewItem, type OrderedMediaPreviewSurface } from "./orderedMediaPreviewTransaction"; -import type { ComfyApp, ComfyImageResult } from "./types"; +import { OrderedMediaDetailSelection } from "./orderedMediaDetailSelection"; +import type { + CubeFaceProjection, + OrderedMediaPreviewActionOptions +} from "./orderedMediaPreviewActionTypes"; +import { + elementSlot, + imageArea, + imageSlot, + indexedSlots, + type NativeImageRect, + previewItems, + unionImageRects, + validSlot +} from "./orderedMediaPreviewGeometry"; +import type { ComfyImageResult } from "./types"; import { comfyImageReferenceKey, comfyImageSourceKey } from "./comfyImageReference"; -type NativeImageRect = readonly [number, number, number, number]; const CUBE_FACE_PROJECTION_SYMBOL = Symbol.for( "sugarcubes.cube-face-projection.v1" ); -interface CubeFaceProjection { - readonly container: HTMLElement; - readonly projectRect: (rect: NativeImageRect) => NativePreviewSlot; -} - -interface ActionPreviewNode { - readonly id?: string | number; - readonly pos?: readonly [number, number]; - readonly size?: readonly [number, number]; - readonly flags?: { readonly collapsed?: boolean }; - widgets?: ActionPreviewWidget[]; - imageIndex?: number | null; - imageRects?: NativeImageRect[]; - imgs?: HTMLImageElement[]; - graph?: { - _rootGraph?: object; - setDirtyCanvas?: (foreground: boolean, background: boolean) => void; - }; -} - -interface ActionPreviewWidget { - readonly y?: number; - readonly computedHeight?: number; - readonly options?: { readonly canvasOnly?: boolean }; -} - -export interface OrderedMediaPreviewActionOptions { - readonly app: ComfyApp; - readonly node: ActionPreviewNode; - readonly preview: NativeNodePreview; - readonly itemLabel: string; - readonly getFiles: () => string[]; - readonly moveEarlier: (index: number) => void; - readonly moveLater: (index: number) => void; - readonly remove: (index: number) => void; -} +export type { OrderedMediaPreviewActionOptions } from "./orderedMediaPreviewActionTypes"; /** Position deterministic list actions over Comfy's native preview cells. */ export class OrderedMediaPreviewActions { private readonly affordances: OrderedMediaPreviewAffordances; private readonly transaction: OrderedMediaPreviewTransaction; + private readonly detailSelection: OrderedMediaDetailSelection; private readonly unsubscribePreview: () => void; private readonly unsubscribeLifecycle: () => void; private lastCanvasPreviewRect: NativeImageRect | null = null; - private pendingDetailIndex: number | null = null; - private detailRestoreFrame: number | null = null; constructor(private readonly options: OrderedMediaPreviewActionOptions) { this.transaction = new OrderedMediaPreviewTransaction({ @@ -83,20 +59,20 @@ export class OrderedMediaPreviewActions { const moveEarlier = (index: number): void => { const destination = index - 1; this.transaction.move(index, destination); - this.followMovedDetail(index, destination); + this.detailSelection.followMoved(index, destination); options.moveEarlier(index); }; const moveLater = (index: number): void => { const destination = index + 1; this.transaction.move(index, destination); - this.followMovedDetail(index, destination); + this.detailSelection.followMoved(index, destination); options.moveLater(index); }; const remove = (index: number): void => { const removedDetail = this.selectedItemIndex() === index; this.transaction.remove(index); options.remove(index); - this.followRemovedDetail(index, removedDetail); + this.detailSelection.followRemoved(index, removedDetail); }; const actionOptions = { getItemCount: () => this.itemCount(), @@ -109,10 +85,30 @@ export class OrderedMediaPreviewActions { getSlots: () => this.nativeActionSlots(), ...actionOptions }); + this.detailSelection = new OrderedMediaDetailSelection({ + selectedIndex: () => this.selectedItemIndex(), + canvasIndex: () => this.options.node.imageIndex, + setCanvasIndex: (index) => { + this.options.node.imageIndex = index; + }, + itemCount: () => this.itemCount(), + domSelectedIndex: () => { + const selectedIndex = this.domDetailButtons().findIndex( + (button) => button.getAttribute("aria-current") === "true" + ); + return selectedIndex >= 0 ? selectedIndex : null; + }, + selectDomIndex: (index) => { + this.domDetailButtons()[index]?.click(); + }, + refreshActions: () => { + this.affordances.refresh(); + } + }); this.unsubscribePreview = options.preview.subscribe(() => { this.affordances.refresh(); this.transaction.authoritativePublished(); - this.restorePendingDetail(); + this.detailSelection.restorePending(); }); this.unsubscribeLifecycle = subscribeNativePreviewLifecycle(() => { this.affordances.refresh(); @@ -125,10 +121,7 @@ export class OrderedMediaPreviewActions { this.unsubscribeLifecycle(); this.transaction.dispose(); this.affordances.dispose(); - if (this.detailRestoreFrame !== null) { - cancelAnimationFrame(this.detailRestoreFrame); - this.detailRestoreFrame = null; - } + this.detailSelection.dispose(); } private itemCount(): number { @@ -492,64 +485,6 @@ export class OrderedMediaPreviewActions { : null; } - /** Keep the moved item selected across both native renderer state models. */ - private followMovedDetail(index: number, destination: number): void { - if (this.selectedItemIndex() !== index) return; - this.pendingDetailIndex = destination; - if (this.options.node.imageIndex === index) { - this.options.node.imageIndex = destination; - } else { - this.domDetailButtons()[destination]?.click(); - } - this.affordances.refresh(); - } - - /** Select the nearest remaining item after removing an inspected item. */ - private followRemovedDetail(index: number, removedDetail: boolean): void { - if (!removedDetail) return; - const remaining = this.itemCount(); - const destination = remaining > 0 ? Math.min(index, remaining - 1) : null; - this.pendingDetailIndex = destination; - if (this.options.node.imageIndex === index) { - this.options.node.imageIndex = destination; - } else if (destination !== null) { - this.domDetailButtons()[destination]?.click(); - } - this.affordances.refresh(); - } - - /** Re-enter Nodes 2.0 detail mode after output publication resets its grid. */ - private restorePendingDetail(): void { - if (this.pendingDetailIndex === null || this.detailRestoreFrame !== null) { - return; - } - let stableFrames = 0; - let remainingFrames = 12; - const restore = (): void => { - this.detailRestoreFrame = null; - const destination = this.pendingDetailIndex; - if (destination === null) return; - const buttons = this.domDetailButtons(); - const selected = buttons.findIndex( - (button) => button.getAttribute("aria-current") === "true" - ); - if (selected === destination) { - stableFrames += 1; - } else { - stableFrames = 0; - buttons[destination]?.click(); - } - remainingFrames -= 1; - if (stableFrames >= 3 || remainingFrames <= 0) { - this.pendingDetailIndex = null; - this.affordances.refresh(); - return; - } - this.detailRestoreFrame = requestAnimationFrame(restore); - }; - this.detailRestoreFrame = requestAnimationFrame(restore); - } - /** Return Comfy's ordered detail navigation controls for this node. */ private domDetailButtons(): HTMLButtonElement[] { const root = this.domRoot(); @@ -566,61 +501,3 @@ export class OrderedMediaPreviewActions { return nodeId === undefined ? undefined : String(nodeId); } } - -function previewItems(images: HTMLImageElement[]): OrderedMediaPreviewItem[] { - return images.map((image) => ({ - sourceUrl: image.currentSrc || image.src, - image - })); -} - -function imageSlot(image: HTMLImageElement): () => NativePreviewSlot { - return () => elementSlot(image); -} - -function elementSlot(element: Element): NativePreviewSlot { - const rect = element.getBoundingClientRect(); - return { - left: rect.left, - top: rect.top, - width: rect.width, - height: rect.height - }; -} - -/** Reject malformed cross-extension projection geometry. */ -function validSlot(slot: NativePreviewSlot): boolean { - return ( - Number.isFinite(slot.left) && - Number.isFinite(slot.top) && - Number.isFinite(slot.width) && - Number.isFinite(slot.height) && - slot.width >= 0 && - slot.height >= 0 - ); -} - -/** Measure a rendered image when native detail contains multiple candidates. */ -function imageArea(image: HTMLImageElement): number { - const rect = image.getBoundingClientRect(); - return rect.width * rect.height; -} - -/** Attach authoritative list positions to native preview footprints. */ -function indexedSlots( - slots: NativePreviewSlot[], - container: HTMLElement | null = null -): NativePreviewActionSlot[] { - return slots.map((bounds, itemIndex) => - container ? { itemIndex, bounds, container } : { itemIndex, bounds } - ); -} - -/** Return the smallest node-local rectangle containing every grid cell. */ -function unionImageRects(rects: NativeImageRect[]): NativeImageRect { - const left = Math.min(...rects.map(([x]) => x)); - const top = Math.min(...rects.map(([, y]) => y)); - const right = Math.max(...rects.map(([x, , width]) => x + width)); - const bottom = Math.max(...rects.map(([, y, , height]) => y + height)); - return [left, top, right - left, bottom - top]; -} diff --git a/web/src/orderedMediaPreviewGeometry.ts b/web/src/orderedMediaPreviewGeometry.ts new file mode 100644 index 0000000..cd3e497 --- /dev/null +++ b/web/src/orderedMediaPreviewGeometry.ts @@ -0,0 +1,75 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +/** Derive ordered-media preview items and renderer-neutral geometry. */ + +import type { + NativePreviewActionSlot, + NativePreviewSlot +} from "./orderedMediaPreviewAffordances"; +import type { OrderedMediaPreviewItem } from "./orderedMediaPreviewTransaction"; + +export type NativeImageRect = readonly [number, number, number, number]; + +export function previewItems( + images: HTMLImageElement[] +): OrderedMediaPreviewItem[] { + return images.map((image) => ({ + sourceUrl: image.currentSrc || image.src, + image + })); +} + +export function imageSlot( + image: HTMLImageElement +): () => NativePreviewSlot { + return () => elementSlot(image); +} + +export function elementSlot(element: Element): NativePreviewSlot { + const rect = element.getBoundingClientRect(); + return { + left: rect.left, + top: rect.top, + width: rect.width, + height: rect.height + }; +} + +/** Reject malformed cross-extension projection geometry. */ +export function validSlot(slot: NativePreviewSlot): boolean { + return ( + Number.isFinite(slot.left) && + Number.isFinite(slot.top) && + Number.isFinite(slot.width) && + Number.isFinite(slot.height) && + slot.width >= 0 && + slot.height >= 0 + ); +} + +/** Measure a rendered image when native detail contains multiple candidates. */ +export function imageArea(image: HTMLImageElement): number { + const rect = image.getBoundingClientRect(); + return rect.width * rect.height; +} + +/** Attach authoritative list positions to native preview footprints. */ +export function indexedSlots( + slots: NativePreviewSlot[], + container: HTMLElement | null = null +): NativePreviewActionSlot[] { + return slots.map((bounds, itemIndex) => + container ? { itemIndex, bounds, container } : { itemIndex, bounds } + ); +} + +/** Return the smallest node-local rectangle containing every grid cell. */ +export function unionImageRects(rects: NativeImageRect[]): NativeImageRect { + const left = Math.min(...rects.map(([x]) => x)); + const top = Math.min(...rects.map(([, y]) => y)); + const right = Math.max(...rects.map(([x, , width]) => x + width)); + const bottom = Math.max(...rects.map(([, y, , height]) => y + height)); + return [left, top, right - left, bottom - top]; +} diff --git a/web/src/quantCacheApi.ts b/web/src/quantCacheApi.ts new file mode 100644 index 0000000..bd5c889 --- /dev/null +++ b/web/src/quantCacheApi.ts @@ -0,0 +1,90 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +/** Quantized-model cache transport and payload validation. */ + +import { backendErrorMessage, type FetchLike } from "./apiTransport"; + +export interface QuantCacheStatus { + path: string; + usage_bytes: number; + limit_bytes: number; + artifact_count: number; + active_artifact_count: number; + removed_artifacts?: number; + removed_bytes?: number; +} + +const QUANT_CACHE_ROUTE = "/simple-syrup/quant-cache"; + +export async function getQuantCacheStatus( + fetchImpl: FetchLike = fetch +): Promise { + return requestQuantCache(fetchImpl); +} + +export async function clearQuantCache( + fetchImpl: FetchLike = fetch +): Promise { + return requestQuantCache(fetchImpl, "DELETE"); +} + +export async function enforceQuantCacheLimit( + fetchImpl: FetchLike = fetch +): Promise { + return requestQuantCache(fetchImpl, "POST"); +} + +export function parseQuantCacheStatus(payload: unknown): QuantCacheStatus { + if (!isQuantCacheStatusPayload(payload)) { + throw new Error( + "SimpleSyrup quant cache status is invalid. Expected path, byte usage, limit, and artifact counts." + ); + } + return { ...payload }; +} + +async function requestQuantCache( + fetchImpl: FetchLike, + method?: "POST" | "DELETE" +): Promise { + const response = + method === undefined + ? await fetchImpl(QUANT_CACHE_ROUTE) + : await fetchImpl(QUANT_CACHE_ROUTE, { method }); + if (!response.ok) { + const fallback = + method === "DELETE" + ? `Could not clear quant cache. Backend returned ${String(response.status)}.` + : method === "POST" + ? `Could not enforce quant cache limit. Backend returned ${String(response.status)}.` + : `Could not load quant cache status. Backend returned ${String(response.status)}.`; + throw new Error( + await backendErrorMessage(response, fallback) + ); + } + return parseQuantCacheStatus(await response.json()); +} + +function isQuantCacheStatusPayload( + payload: unknown +): payload is QuantCacheStatus { + if (typeof payload !== "object" || payload === null) return false; + const candidate = payload as Partial; + return ( + typeof candidate.path === "string" && + isNonNegativeInteger(candidate.usage_bytes) && + isNonNegativeInteger(candidate.limit_bytes) && + isNonNegativeInteger(candidate.artifact_count) && + isNonNegativeInteger(candidate.active_artifact_count) && + (candidate.removed_artifacts === undefined || + isNonNegativeInteger(candidate.removed_artifacts)) && + (candidate.removed_bytes === undefined || + isNonNegativeInteger(candidate.removed_bytes)) + ); +} + +function isNonNegativeInteger(value: unknown): value is number { + return typeof value === "number" && Number.isInteger(value) && value >= 0; +} diff --git a/web/src/quantCacheSetting.ts b/web/src/quantCacheSetting.ts index ed7809b..44a0f5d 100644 --- a/web/src/quantCacheSetting.ts +++ b/web/src/quantCacheSetting.ts @@ -2,7 +2,7 @@ // Copyright (C) 2026 Artificial Sweetener and contributors // SPDX-License-Identifier: AGPL-3.0-or-later -import type { QuantCacheStatus } from "./api"; +import type { QuantCacheStatus } from "./quantCacheApi"; import type { ComfyApp, Logger } from "./types"; import type { GeneralSettingsContext } from "./downloadableModelsSetting"; import { diff --git a/web/src/refresh.ts b/web/src/refresh.ts index 9857f23..454e2f4 100644 --- a/web/src/refresh.ts +++ b/web/src/refresh.ts @@ -2,7 +2,7 @@ // Copyright (C) 2026 Artificial Sweetener and contributors // SPDX-License-Identifier: AGPL-3.0-or-later -import { refreshExternalLLMModels } from "./api"; +import { refreshExternalLLMModels } from "./externalLlmApi"; import type { ComfyApp, Logger } from "./types"; const REFRESH_WRAPPED = Symbol.for("SimpleSyrup.ExternalLLM.RefreshWrapped"); diff --git a/web/src/settingsRegistration.ts b/web/src/settingsRegistration.ts index 8b5dc27..7b55fc5 100644 --- a/web/src/settingsRegistration.ts +++ b/web/src/settingsRegistration.ts @@ -3,20 +3,22 @@ // SPDX-License-Identifier: AGPL-3.0-or-later import { - clearQuantCache, - enforceQuantCacheLimit, - getExternalLLMSettings, - getQuantCacheStatus, getSettings, - saveExternalLLMApiKey, - saveExternalLLMSettings, saveSettings } from "./api"; -import type { - ExternalLLMSettings, - QuantCacheStatus, - SimpleSyrupSettings -} from "./api"; +import type { SimpleSyrupSettings } from "./api"; +import { + getExternalLLMSettings, + saveExternalLLMApiKey, + saveExternalLLMSettings +} from "./externalLlmApi"; +import type { ExternalLLMSettings } from "./externalLlmApi"; +import { + clearQuantCache, + enforceQuantCacheLimit, + getQuantCacheStatus +} from "./quantCacheApi"; +import type { QuantCacheStatus } from "./quantCacheApi"; import { registerDownloadableModelsSetting, type GeneralSettingsContext diff --git a/web/tests/api.test.ts b/web/tests/comfy_integration/api.test.ts similarity index 96% rename from web/tests/api.test.ts rename to web/tests/comfy_integration/api.test.ts index 06ea864..bdc293f 100644 --- a/web/tests/api.test.ts +++ b/web/tests/comfy_integration/api.test.ts @@ -5,24 +5,30 @@ import { describe, expect, it, vi } from "vitest"; import { - clearQuantCache, + getSettings, + parseSettings, + saveSettings +} from "../../src/api"; +import type { FetchLike } from "../../src/apiTransport"; +import { deleteExternalLLMApiKey, getExternalLLMSettings, - getMaskBatchPreview, - getQuantCacheStatus, - getSettings, - enforceQuantCacheLimit, parseExternalLLMSettings, - parseMaskBatchPreview, - parseQuantCacheStatus, - parseSettings, refreshExternalLLMModels, saveExternalLLMApiKey, - saveExternalLLMSettings, - saveSettings -} from "../src/api"; -import type { FetchLike } from "../src/api"; -import { createJsonResponse } from "./testUtils"; + saveExternalLLMSettings +} from "../../src/externalLlmApi"; +import { + getMaskBatchPreview, + parseMaskBatchPreview +} from "../../src/maskBatchPreviewApi"; +import { + clearQuantCache, + enforceQuantCacheLimit, + getQuantCacheStatus, + parseQuantCacheStatus +} from "../../src/quantCacheApi"; +import { createJsonResponse } from "../support/testUtils"; describe("settings API", () => { it("loads SimpleSyrup settings from the backend route", async () => { diff --git a/web/tests/refresh.test.ts b/web/tests/comfy_integration/refresh.test.ts similarity index 94% rename from web/tests/refresh.test.ts rename to web/tests/comfy_integration/refresh.test.ts index 2e8b1dc..bf9e8f4 100644 --- a/web/tests/refresh.test.ts +++ b/web/tests/comfy_integration/refresh.test.ts @@ -4,8 +4,8 @@ import { describe, expect, it, vi } from "vitest"; -import { registerExternalLLMRefreshHook } from "../src/refresh"; -import { createFakeComfyApp } from "./testUtils"; +import { registerExternalLLMRefreshHook } from "../../src/refresh"; +import { createFakeComfyApp } from "../support/testUtils"; describe("external LLM refresh hook", () => { it("refreshes external models before Comfy refreshes node definitions", async () => { diff --git a/web/tests/imageListUpload.test.ts b/web/tests/media/imageListUpload.test.ts similarity index 82% rename from web/tests/imageListUpload.test.ts rename to web/tests/media/imageListUpload.test.ts index a087ab3..79f3fec 100644 --- a/web/tests/imageListUpload.test.ts +++ b/web/tests/media/imageListUpload.test.ts @@ -4,9 +4,9 @@ import { describe, expect, it } from "vitest"; -import { inputImageReference } from "../src/comfyImageUrl"; -import { registerImageListUpload } from "../src/imageListUpload"; -import type { ComfyApp, ComfyExtension } from "../src/types"; +import { inputImageReference } from "../../src/comfyImageUrl"; +import { registerImageListUpload } from "../../src/imageListUpload"; +import type { ComfyApp, ComfyExtension } from "../../src/types"; describe("Load Image List frontend integration", () => { it("registers its native ordered-preview extension", () => { diff --git a/web/tests/imageViewport.test.ts b/web/tests/media/imageViewport.test.ts similarity index 92% rename from web/tests/imageViewport.test.ts rename to web/tests/media/imageViewport.test.ts index eed3b30..de0837a 100644 --- a/web/tests/imageViewport.test.ts +++ b/web/tests/media/imageViewport.test.ts @@ -4,7 +4,7 @@ import { describe, expect, it } from "vitest"; -import { ImageViewport } from "../src/imageViewport"; +import { ImageViewport } from "../../src/imageViewport"; describe("ImageViewport", () => { it("maps displayed pointer and crop geometry to the source image", () => { diff --git a/web/tests/interactiveInspector.test.ts b/web/tests/media/interactiveInspector.test.ts similarity index 98% rename from web/tests/interactiveInspector.test.ts rename to web/tests/media/interactiveInspector.test.ts index 5521d3c..27ebfb6 100644 --- a/web/tests/interactiveInspector.test.ts +++ b/web/tests/media/interactiveInspector.test.ts @@ -8,7 +8,7 @@ import { InteractiveInspectorController, type AsyncInspectorView, type PreparedInspectorState -} from "../src/interactiveInspector"; +} from "../../src/interactiveInspector"; describe("InteractiveInspectorController", () => { it("discards stale prepared documents and commits only the latest", async () => { diff --git a/web/tests/nativeNodePreview.test.ts b/web/tests/media/nativeNodePreview.test.ts similarity index 98% rename from web/tests/nativeNodePreview.test.ts rename to web/tests/media/nativeNodePreview.test.ts index 2a8e8d5..a7f28d7 100644 --- a/web/tests/nativeNodePreview.test.ts +++ b/web/tests/media/nativeNodePreview.test.ts @@ -4,7 +4,7 @@ import { describe, expect, it, vi } from "vitest"; -import { NativeNodePreview } from "../src/nativeNodePreview"; +import { NativeNodePreview } from "../../src/nativeNodePreview"; import type { ComfyApi, ComfyApp, @@ -12,7 +12,7 @@ import type { ComfySetting, ComfySettingDefinition, SettingValue -} from "../src/types"; +} from "../../src/types"; describe("NativeNodePreview", () => { it("publishes execution-shaped output through Comfy's native event path", () => { diff --git a/web/tests/nativePreviewNavigator.test.ts b/web/tests/media/nativePreviewNavigator.test.ts similarity index 97% rename from web/tests/nativePreviewNavigator.test.ts rename to web/tests/media/nativePreviewNavigator.test.ts index c9de2a3..31d54f1 100644 --- a/web/tests/nativePreviewNavigator.test.ts +++ b/web/tests/media/nativePreviewNavigator.test.ts @@ -4,9 +4,9 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { NativePreviewNavigator } from "../src/nativePreviewNavigator"; -import type { ComfyApp } from "../src/types"; -import { createFakeComfyApp } from "./testUtils"; +import { NativePreviewNavigator } from "../../src/nativePreviewNavigator"; +import type { ComfyApp } from "../../src/types"; +import { createFakeComfyApp } from "../support/testUtils"; describe("NativePreviewNavigator", () => { afterEach(() => { diff --git a/web/tests/media/orderedMediaNode.test.ts b/web/tests/media/orderedMediaNode.test.ts new file mode 100644 index 0000000..24eb0f7 --- /dev/null +++ b/web/tests/media/orderedMediaNode.test.ts @@ -0,0 +1,186 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { afterEach, describe, expect, it, vi } from "vitest"; + +import { + configureOrderedMediaNode, + registerOrderedMediaNode +} from "../../src/orderedMediaNode"; +import type { ComfyExtension } from "../../src/types"; +import { + CONFIG, + configured, + createFixture, + references, + resetOrderedMediaFixtures, + selectFiles, + widget +} from "./orderedMediaNodeTestSupport"; + +afterEach(resetOrderedMediaFixtures); + +describe("ordered-media node integration", () => { + it.each(["Nodes 1.0", "Nodes 2.0"])( + "publishes through Comfy's native preview under %s", + async () => { + const fixture = createFixture(["one.png"]); + configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, CONFIG); + + expect(fixture.imageWidget.hidden).toBe(true); + expect(fixture.uploadWidget.hidden).toBe(true); + expect(fixture.imageWidget.computeSize?.()).toEqual([0, -4]); + expect( + fixture.node.widgets.some((candidate) => + [ + "simple_syrup_selected_media", + "simple_syrup_move_earlier", + "simple_syrup_move_later", + "simple_syrup_remove_media" + ].includes(candidate.name) + ) + ).toBe(false); + expect(widget(fixture, "simple_syrup_replace_media").label).toBe( + "Replace images..." + ); + expect(widget(fixture, "simple_syrup_add_media").label).toBe( + "Add images..." + ); + expect("addDOMWidget" in fixture.node).toBe(false); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png") + ); + }); + expect(fixture.executed).toHaveBeenCalledWith( + expect.objectContaining({ node: "7", display_node: "7" }) + ); + } + ); + + it("appends native multi-upload results and preserves duplicates", async () => { + const fixture = configured(["one.png", "same.png"]); + + widget(fixture, "simple_syrup_add_media").callback?.(); + fixture.imageWidget.value = "two.png"; + fixture.imageWidget.callback?.("two.png"); + fixture.imageWidget.value = ["two.png", "same.png"]; + fixture.imageWidget.callback?.(["two.png", "same.png"]); + + expect(fixture.nativeUpload).toHaveBeenCalledOnce(); + expect(fixture.imageWidget.value).toEqual([ + "one.png", + "same.png", + "two.png", + "same.png" + ]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png", "same.png", "two.png", "same.png") + ); + }); + }); + + it("treats pasted and dropped files as append operations", () => { + const fixture = configured(["existing.png"]); + + fixture.node.onDragDrop(); + selectFiles(fixture, ["dropped-a.png", "dropped-b.png"]); + expect(fixture.imageWidget.value).toEqual([ + "existing.png", + "dropped-a.png", + "dropped-b.png" + ]); + + fixture.node.pasteFiles(); + selectFiles(fixture, ["pasted.png"]); + expect(fixture.imageWidget.value).toEqual([ + "existing.png", + "dropped-a.png", + "dropped-b.png", + "pasted.png" + ]); + }); + + it("normalizes scalar workflows and survives Nodes 2.0 reactive assignments", async () => { + const fixture = configured("saved.png", true); + + const onGraphConfigured = fixture.node.onGraphConfigured; + if (!onGraphConfigured) throw new Error("Expected graph configuration callback."); + onGraphConfigured(); + + expect(fixture.imageWidget.value).toEqual(["saved.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("saved.png") + ); + }); + }); + + it("retains array-valued media when the legacy graph lifecycle clears the widget", async () => { + const fixture = createFixture(["one.png", "two.png"]); + fixture.node.onGraphConfigured = () => { + fixture.imageWidget.value = undefined; + }; + configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, CONFIG); + fixture.node.onGraphConfigured(); + + expect(fixture.imageWidget.value).toEqual(["one.png", "two.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png", "two.png") + ); + }); + }); + + it("restores persisted widget values when Comfy clears the live media widget", async () => { + const fixture = createFixture(undefined); + fixture.node.widgets_values = [ + ["one.png", "two.png"], + "image", + "ordered_media", + "ordered_media" + ]; + configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, CONFIG); + fixture.imageWidget.value = undefined; + + fixture.node.onGraphConfigured?.(); + + expect(fixture.imageWidget.value).toEqual(["one.png", "two.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png", "two.png") + ); + }); + }); + + it("restores native handlers and clears preview on node removal", () => { + const fixture = createFixture(["one.png"]); + const paste = fixture.node.pasteFiles; + const drop = fixture.node.onDragDrop; + configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, CONFIG); + + fixture.node.onRemoved(); + + expect(fixture.node.pasteFiles).toBe(paste); + expect(fixture.node.onDragDrop).toBe(drop); + expect(fixture.originalRemoved).toHaveBeenCalledOnce(); + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual([]); + }); + + it("registers a guarded extension and ignores unrelated nodes", async () => { + const fixture = createFixture([]); + let extension: ComfyExtension | undefined; + fixture.app.registerExtension = (value: ComfyExtension) => { + extension = value; + }; + registerOrderedMediaNode(fixture.app, fixture.api, "test.extension", CONFIG); + fixture.node.constructor.comfyClass = "Other.Node"; + + await extension?.nodeCreated?.(fixture.node); + + expect(extension?.name).toBe("test.extension"); + expect(fixture.node.addWidget).not.toHaveBeenCalled(); + }); +}); diff --git a/web/tests/media/orderedMediaNodePreview.test.ts b/web/tests/media/orderedMediaNodePreview.test.ts new file mode 100644 index 0000000..5ceb45a --- /dev/null +++ b/web/tests/media/orderedMediaNodePreview.test.ts @@ -0,0 +1,292 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { afterEach, describe, expect, it, vi } from "vitest"; + +import { configureOrderedMediaNode } from "../../src/orderedMediaNode"; +import type { ComfyNodeExecutionOutput } from "../../src/types"; +import { + CONFIG, + configured, + createFixture, + deferred, + imageFiles, + mockRect, + nativePreviewButton, + nativePreviewImage, + references, + resetOrderedMediaFixtures, + visiblePreviewFiles +} from "./orderedMediaNodeTestSupport"; + +afterEach(resetOrderedMediaFixtures); + +describe("ordered-media native preview integration", () => { + it("attaches actions when a multi-image native gallery mounts later", async () => { + const fixture = configured(["one.png", "two.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(2); + }); + expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(0); + const root = document.createElement("section"); + root.dataset.nodeId = "7"; + root.append( + nativePreviewButton("one.png", 0), + nativePreviewButton("two.png", 90) + ); + + document.body.append(root); + + await vi.waitFor(() => { + expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(2); + }); + }); + + it("moves and removes exact positions through thumbnail actions", async () => { + const fixture = configured(["same.png", "middle.png", "same.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(3); + }); + const root = document.createElement("section"); + root.dataset.nodeId = "7"; + root.append( + nativePreviewImage("same.png", 0), + nativePreviewImage("middle.png", 90), + nativePreviewImage("same.png", 180) + ); + document.body.append(root); + await vi.waitFor(() => { + expect(document.querySelectorAll("[data-ss-media-move-earlier]")).toHaveLength(3); + }); + document + .querySelectorAll("[data-ss-media-move-earlier]")[2] + ?.click(); + + expect(fixture.imageWidget.value).toEqual([ + "same.png", + "same.png", + "middle.png" + ]); + expect(visiblePreviewFiles(root)).toEqual([ + "same.png", + "same.png", + "middle.png" + ]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("same.png", "same.png", "middle.png") + ); + }); + expect(fixture.onWidgetChanged).toHaveBeenCalledWith( + "image", + ["same.png", "same.png", "middle.png"], + ["same.png", "middle.png", "same.png"], + fixture.imageWidget + ); + + document + .querySelectorAll("[data-ss-media-remove]")[1] + ?.click(); + expect(fixture.imageWidget.value).toEqual(["same.png", "middle.png"]); + expect(visiblePreviewFiles(root)).toEqual(["same.png", "middle.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("same.png", "middle.png") + ); + }); + }); + + it("keeps exactly the native cells in visual order during a slow rebuild", async () => { + const delayed = deferred(); + let previewCalls = 0; + const config = { + ...CONFIG, + preview: (files: string[]) => { + previewCalls += 1; + return previewCalls === 1 + ? Promise.resolve({ images: references(...files) }) + : delayed.promise; + } + }; + const fixture = createFixture(["one.png", "two.png", "three.png"]); + configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, config); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png", "two.png", "three.png") + ); + }); + const root = document.createElement("section"); + root.dataset.nodeId = "7"; + root.append( + nativePreviewImage("one.png", 0), + nativePreviewImage("two.png", 90), + nativePreviewImage("three.png", 180) + ); + document.body.append(root); + await vi.waitFor(() => { + expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(3); + }); + + document + .querySelectorAll("[data-ss-media-move-later]")[0] + ?.click(); + + expect(fixture.imageWidget.value).toEqual([ + "two.png", + "one.png", + "three.png" + ]); + expect(visiblePreviewFiles(root)).toEqual([ + "two.png", + "one.png", + "three.png" + ]); + expect(root.querySelectorAll("img")).toHaveLength(3); + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png", "two.png", "three.png") + ); + await new Promise((resolve) => setTimeout(resolve, 50)); + expect(visiblePreviewFiles(root)).toEqual([ + "two.png", + "one.png", + "three.png" + ]); + + delayed.resolve({ images: references("two.png", "one.png", "three.png") }); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("two.png", "one.png", "three.png") + ); + }); + expect(visiblePreviewFiles(root)).toEqual([ + "two.png", + "one.png", + "three.png" + ]); + + root.replaceChildren( + nativePreviewImage("two.png", 0), + nativePreviewImage("one.png", 90), + nativePreviewImage("three.png", 180) + ); + expect(visiblePreviewFiles(root)).toEqual([ + "two.png", + "one.png", + "three.png" + ]); + expect(root.querySelectorAll("img")).toHaveLength(3); + }); + + it("keeps the native visual and serialized order when rebuilding fails", async () => { + const logger = { warn: vi.fn() }; + let previewCalls = 0; + const config = { + ...CONFIG, + preview: (files: string[]) => { + previewCalls += 1; + return previewCalls === 1 + ? Promise.resolve({ images: references(...files) }) + : Promise.reject(new Error("preview unavailable")); + } + }; + const fixture = createFixture(["one.png", "two.png"]); + configureOrderedMediaNode( + fixture.node, + fixture.app, + fixture.api, + config, + logger + ); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png", "two.png") + ); + }); + const root = document.createElement("section"); + root.dataset.nodeId = "7"; + root.append( + nativePreviewImage("one.png", 0), + nativePreviewImage("two.png", 90) + ); + document.body.append(root); + await vi.waitFor(() => { + expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(2); + }); + + document + .querySelectorAll("[data-ss-media-move-later]")[0] + ?.click(); + + await vi.waitFor(() => { + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("preview unavailable"), + expect.any(Error) + ); + }); + expect(fixture.imageWidget.value).toEqual(["two.png", "one.png"]); + expect(visiblePreviewFiles(root)).toEqual(["two.png", "one.png"]); + expect(root.querySelectorAll("img")).toHaveLength(2); + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("one.png", "two.png") + ); + }); + + it("preserves native thumbnail inspection clicks in Nodes 2.0", async () => { + const fixture = configured(["one.png", "two.png", "three.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(3); + }); + const root = document.createElement("section"); + root.dataset.nodeId = "7"; + const images = ["one.png", "two.png", "three.png"].map( + (filename, index) => nativePreviewImage(filename, index * 90) + ); + root.append(...images); + document.body.append(root); + const clicked = vi.fn(); + images[1]?.addEventListener("click", clicked); + + images[1]?.dispatchEvent(new MouseEvent("click", { bubbles: true })); + expect(clicked).toHaveBeenCalledOnce(); + }); + + it("moves Nodes 1.0 native canvas cells through imageRects", async () => { + const fixture = createFixture(["one.png", "two.png"]); + const canvas = document.createElement("canvas"); + mockRect(canvas, 0, 0, 500, 500); + document.body.append(canvas); + fixture.app.canvas = { + canvas, + convertEventToCanvasOffset: (event) => [event.clientX, event.clientY], + convertOffsetToCanvas: (position) => [...position] + }; + Object.assign(fixture.node, { + pos: [10, 100], + imageIndex: null, + imageRects: [ + [0, 50, 80, 80], + [80, 50, 80, 80] + ], + imgs: [nativePreviewImage("one.png", 0), nativePreviewImage("two.png", 90)] + }); + configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, CONFIG); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(2); + }); + await vi.waitFor(() => { + expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(2); + }); + document + .querySelectorAll("[data-ss-media-move-later]")[0] + ?.click(); + + expect(fixture.imageWidget.value).toEqual(["two.png", "one.png"]); + expect(imageFiles(fixture.node.imgs)).toEqual(["two.png", "one.png"]); + await vi.waitFor(() => { + expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( + references("two.png", "one.png") + ); + }); + }); +}); diff --git a/web/tests/media/orderedMediaNodeTestSupport.ts b/web/tests/media/orderedMediaNodeTestSupport.ts new file mode 100644 index 0000000..32f17e3 --- /dev/null +++ b/web/tests/media/orderedMediaNodeTestSupport.ts @@ -0,0 +1,233 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +/** Provide focused ordered-media node fixtures shared across capability tests. */ + +import { vi } from "vitest"; + +import { configureOrderedMediaNode } from "../../src/orderedMediaNode"; +import type { + ComfyApi, + ComfyApp, + ComfyImageResult +} from "../../src/types"; + +export const CONFIG = { + nodeId: "SimpleSyrup.LoadImageList", + labels: { + singular: "image", + plural: "images", + replace: "Replace images...", + add: "Add images..." + }, + preview: (files: string[]) => + Promise.resolve({ + images: references(...files), + animated: files.map(() => false) + }) +}; + +interface TestWidget { + name: string; + value: unknown; + type?: string; + label?: string; + hidden?: boolean; + disabled?: boolean; + callback?: (value?: unknown) => void; + computeSize?: (width?: number) => [number, number]; + options?: { + canvasOnly?: boolean; + serialize?: boolean; + tooltip?: string; + hidden?: boolean; + }; +} + +const activeNodes: Array<{ onRemoved: () => unknown }> = []; + +export function resetOrderedMediaFixtures(): void { + for (const node of activeNodes.splice(0)) node.onRemoved(); + document.body.replaceChildren(); +} + +export function configured(value: unknown, reactive = false) { + const fixture = createFixture(value, reactive); + configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, CONFIG); + return fixture; +} + +export function createFixture(value: unknown, reactive = false) { + const nativeUpload = vi.fn(); + const imageWidget: TestWidget = { + name: "image", + value, + type: "combo", + options: { canvasOnly: true } + }; + if (reactive) { + let stored = value; + Object.defineProperty(imageWidget, "value", { + configurable: true, + get: () => stored, + set: (next: unknown) => { + stored = next; + imageWidget.callback?.(next); + } + }); + } + const uploadWidget: TestWidget = { + name: "upload", + value: "image", + type: "button", + callback: nativeUpload, + options: { serialize: false, canvasOnly: true } + }; + const widgets = [imageWidget, uploadWidget]; + const originalRemoved = vi.fn(); + const onWidgetChanged = vi.fn(); + const node = { + constructor: { comfyClass: "SimpleSyrup.LoadImageList" }, + id: 7, + widgets, + widgets_values: undefined as unknown[] | undefined, + imgs: undefined as HTMLImageElement[] | undefined, + graph: { setDirtyCanvas: vi.fn() }, + pasteFiles: vi.fn(() => true) as (...args: unknown[]) => unknown, + onDragDrop: vi.fn(() => true) as (...args: unknown[]) => unknown, + onRemoved: originalRemoved as (...args: unknown[]) => unknown, + onGraphConfigured: undefined as ((...args: unknown[]) => unknown) | undefined, + onWidgetChanged, + addWidget: vi.fn( + ( + type: "button" | "combo", + name: string, + widgetValue: string, + callback: (value?: unknown) => void, + options?: TestWidget["options"] + ): TestWidget => { + const created: TestWidget = { + type, + name, + value: widgetValue, + callback + }; + if (options) created.options = options; + widgets.push(created); + return created; + } + ) + }; + activeNodes.push(node); + const app = { + nodeOutputs: {}, + ui: { settings: { addSetting: vi.fn() } }, + registerExtension: vi.fn() + } as unknown as ComfyApp; + const events = new EventTarget(); + const api = events as ComfyApi; + const executed = vi.fn(); + events.addEventListener("executed", (event: Event) => { + executed((event as CustomEvent).detail); + }); + return { + api, + app, + executed, + imageWidget, + nativeUpload, + node, + onWidgetChanged, + originalRemoved, + uploadWidget + }; +} + +export function nativePreviewImage( + filename: string, + left: number +): HTMLImageElement { + const image = document.createElement("img"); + image.src = `/api/view?filename=${filename}&subfolder=&type=input&preview=webp`; + Object.defineProperties(image, { + complete: { configurable: true, value: true }, + naturalWidth: { configurable: true, value: 80 } + }); + mockRect(image, left, 0, 80, 80); + return image; +} + +export function nativePreviewButton( + filename: string, + left: number +): HTMLButtonElement { + const button = document.createElement("button"); + button.append(nativePreviewImage(filename, left)); + mockRect(button, left, 0, 80, 80); + return button; +} + +export function visiblePreviewFiles(root: HTMLElement): string[] { + return Array.from(root.querySelectorAll("img")) + .filter((image) => image.style.display !== "none") + .map((image) => new URL(image.src).searchParams.get("filename") ?? ""); +} + +export function imageFiles(images: HTMLImageElement[] | undefined): string[] { + return (images ?? []).map( + (image) => new URL(image.src).searchParams.get("filename") ?? "" + ); +} + +export function mockRect( + element: Element, + left: number, + top: number, + width: number, + height: number +): void { + element.getBoundingClientRect = () => new DOMRect(left, top, width, height); +} + +export function widget( + fixture: ReturnType, + name: string +): TestWidget { + const found = fixture.node.widgets.find((candidate) => candidate.name === name); + if (!found) throw new Error(`Missing test widget ${name}.`); + return found; +} + +export function selectFiles( + fixture: ReturnType, + files: string[] +): void { + fixture.imageWidget.value = files; + fixture.imageWidget.callback?.(files); +} + +export function references(...filenames: string[]): ComfyImageResult[] { + return filenames.map((filename) => ({ + filename, + subfolder: "", + type: "input" + })); +} + +export function deferred(): { + readonly promise: Promise; + readonly resolve: (value: T) => void; +} { + let resolvePromise: ((value: T) => void) | undefined; + const promise = new Promise((resolve) => { + resolvePromise = resolve; + }); + return { + promise, + resolve: (value) => { + if (!resolvePromise) throw new Error("Deferred promise is unavailable."); + resolvePromise(value); + } + }; +} diff --git a/web/tests/orderedMediaPreview.test.ts b/web/tests/media/orderedMediaPreview.test.ts similarity index 96% rename from web/tests/orderedMediaPreview.test.ts rename to web/tests/media/orderedMediaPreview.test.ts index ce4506e..959c4a3 100644 --- a/web/tests/orderedMediaPreview.test.ts +++ b/web/tests/media/orderedMediaPreview.test.ts @@ -4,8 +4,8 @@ import { describe, expect, it, vi } from "vitest"; -import { OrderedMediaPreviewController } from "../src/orderedMediaPreview"; -import type { ComfyNodeExecutionOutput, Logger } from "../src/types"; +import { OrderedMediaPreviewController } from "../../src/orderedMediaPreview"; +import type { ComfyNodeExecutionOutput, Logger } from "../../src/types"; const OUTPUT: ComfyNodeExecutionOutput = { images: [{ filename: "mask.png", subfolder: "", type: "temp" }], diff --git a/web/tests/orderedMediaPreviewActions.test.ts b/web/tests/media/orderedMediaPreviewActions.test.ts similarity index 98% rename from web/tests/orderedMediaPreviewActions.test.ts rename to web/tests/media/orderedMediaPreviewActions.test.ts index ac783a2..aba50af 100644 --- a/web/tests/orderedMediaPreviewActions.test.ts +++ b/web/tests/media/orderedMediaPreviewActions.test.ts @@ -4,9 +4,9 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { NativeNodePreview } from "../src/nativeNodePreview"; -import { OrderedMediaPreviewActions } from "../src/orderedMediaPreviewActions"; -import type { ComfyApi, ComfyApp, ComfyImageResult } from "../src/types"; +import { NativeNodePreview } from "../../src/nativeNodePreview"; +import { OrderedMediaPreviewActions } from "../../src/orderedMediaPreviewActions"; +import type { ComfyApi, ComfyApp, ComfyImageResult } from "../../src/types"; const adapters: OrderedMediaPreviewActions[] = []; diff --git a/web/tests/orderedMediaPreviewAffordances.test.ts b/web/tests/media/orderedMediaPreviewAffordances.test.ts similarity index 98% rename from web/tests/orderedMediaPreviewAffordances.test.ts rename to web/tests/media/orderedMediaPreviewAffordances.test.ts index f99bd45..58aae62 100644 --- a/web/tests/orderedMediaPreviewAffordances.test.ts +++ b/web/tests/media/orderedMediaPreviewAffordances.test.ts @@ -4,7 +4,7 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { OrderedMediaPreviewAffordances } from "../src/orderedMediaPreviewAffordances"; +import { OrderedMediaPreviewAffordances } from "../../src/orderedMediaPreviewAffordances"; let affordances: OrderedMediaPreviewAffordances | undefined; diff --git a/web/tests/orderedMediaPreviewTransaction.test.ts b/web/tests/media/orderedMediaPreviewTransaction.test.ts similarity index 95% rename from web/tests/orderedMediaPreviewTransaction.test.ts rename to web/tests/media/orderedMediaPreviewTransaction.test.ts index 02044de..73e1b13 100644 --- a/web/tests/orderedMediaPreviewTransaction.test.ts +++ b/web/tests/media/orderedMediaPreviewTransaction.test.ts @@ -7,8 +7,8 @@ import { describe, expect, it, vi } from "vitest"; import { OrderedMediaPreviewTransaction, type OrderedMediaPreviewItem -} from "../src/orderedMediaPreviewTransaction"; -import type { NativePreviewSlot } from "../src/orderedMediaPreviewAffordances"; +} from "../../src/orderedMediaPreviewTransaction"; +import type { NativePreviewSlot } from "../../src/orderedMediaPreviewAffordances"; describe("OrderedMediaPreviewTransaction", () => { it("reorders the captured native surface without creating DOM elements", async () => { diff --git a/web/tests/orderedMediaSelection.test.ts b/web/tests/media/orderedMediaSelection.test.ts similarity index 96% rename from web/tests/orderedMediaSelection.test.ts rename to web/tests/media/orderedMediaSelection.test.ts index 392543e..e39070b 100644 --- a/web/tests/orderedMediaSelection.test.ts +++ b/web/tests/media/orderedMediaSelection.test.ts @@ -7,7 +7,7 @@ import { describe, expect, it } from "vitest"; import { normalizeMediaFiles, OrderedMediaSelection -} from "../src/orderedMediaSelection"; +} from "../../src/orderedMediaSelection"; describe("OrderedMediaSelection", () => { it("normalizes scalar and invalid persisted values", () => { diff --git a/web/tests/selectionModel.test.ts b/web/tests/media/selectionModel.test.ts similarity index 94% rename from web/tests/selectionModel.test.ts rename to web/tests/media/selectionModel.test.ts index d03b292..fb8c7f6 100644 --- a/web/tests/selectionModel.test.ts +++ b/web/tests/media/selectionModel.test.ts @@ -4,7 +4,7 @@ import { describe, expect, it, vi } from "vitest"; -import { SelectionModel } from "../src/selectionModel"; +import { SelectionModel } from "../../src/selectionModel"; describe("SelectionModel", () => { it("shares hover and pinned selection across inspector views", () => { diff --git a/web/tests/orderedMediaNode.test.ts b/web/tests/orderedMediaNode.test.ts deleted file mode 100644 index 5921ccb..0000000 --- a/web/tests/orderedMediaNode.test.ts +++ /dev/null @@ -1,699 +0,0 @@ -// SimpleSyrup - workflow-focused ComfyUI extensions for image generation -// Copyright (C) 2026 Artificial Sweetener and contributors -// SPDX-License-Identifier: AGPL-3.0-or-later - -import { afterEach, describe, expect, it, vi } from "vitest"; - -import { - configureOrderedMediaNode, - registerOrderedMediaNode -} from "../src/orderedMediaNode"; -import type { - ComfyApi, - ComfyApp, - ComfyExtension, - ComfyImageResult, - ComfyNodeExecutionOutput -} from "../src/types"; - -const CONFIG = { - nodeId: "SimpleSyrup.LoadImageList", - labels: { - singular: "image", - plural: "images", - replace: "Replace images...", - add: "Add images..." - }, - preview: (files: string[]) => - Promise.resolve({ - images: references(...files), - animated: files.map(() => false) - }) -}; - -const activeNodes: Array<{ onRemoved: () => unknown }> = []; - -afterEach(() => { - for (const node of activeNodes.splice(0)) node.onRemoved(); - document.body.replaceChildren(); -}); - -describe("ordered-media node integration", () => { - it.each(["Nodes 1.0", "Nodes 2.0"])( - "publishes through Comfy's native preview under %s", - async () => { - const fixture = createFixture(["one.png"]); - configureOrderedMediaNode( - fixture.node, - fixture.app, - fixture.api, - CONFIG - ); - - expect(fixture.imageWidget.hidden).toBe(true); - expect(fixture.uploadWidget.hidden).toBe(true); - expect(fixture.imageWidget.computeSize?.()).toEqual([0, -4]); - expect( - fixture.node.widgets.some((candidate) => - [ - "simple_syrup_selected_media", - "simple_syrup_move_earlier", - "simple_syrup_move_later", - "simple_syrup_remove_media" - ].includes(candidate.name) - ) - ).toBe(false); - expect(widget(fixture, "simple_syrup_replace_media").label).toBe( - "Replace images..." - ); - expect(widget(fixture, "simple_syrup_add_media").label).toBe( - "Add images..." - ); - expect("addDOMWidget" in fixture.node).toBe(false); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png") - ); - }); - expect(fixture.executed).toHaveBeenCalledWith( - expect.objectContaining({ node: "7", display_node: "7" }) - ); - } - ); - - it("appends native multi-upload results and preserves duplicates", async () => { - const fixture = configured(["one.png", "same.png"]); - - widget(fixture, "simple_syrup_add_media").callback?.(); - fixture.imageWidget.value = "two.png"; - fixture.imageWidget.callback?.("two.png"); - fixture.imageWidget.value = ["two.png", "same.png"]; - fixture.imageWidget.callback?.(["two.png", "same.png"]); - - expect(fixture.nativeUpload).toHaveBeenCalledOnce(); - expect(fixture.imageWidget.value).toEqual([ - "one.png", - "same.png", - "two.png", - "same.png" - ]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png", "same.png", "two.png", "same.png") - ); - }); - }); - - it("treats pasted and dropped files as append operations", () => { - const fixture = configured(["existing.png"]); - - fixture.node.onDragDrop(); - selectFiles(fixture, ["dropped-a.png", "dropped-b.png"]); - expect(fixture.imageWidget.value).toEqual([ - "existing.png", - "dropped-a.png", - "dropped-b.png" - ]); - - fixture.node.pasteFiles(); - selectFiles(fixture, ["pasted.png"]); - expect(fixture.imageWidget.value).toEqual([ - "existing.png", - "dropped-a.png", - "dropped-b.png", - "pasted.png" - ]); - }); - - it("attaches actions when a multi-image native gallery mounts later", async () => { - const fixture = configured(["one.png", "two.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(2); - }); - expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(0); - const root = document.createElement("section"); - root.dataset.nodeId = "7"; - root.append( - nativePreviewButton("one.png", 0), - nativePreviewButton("two.png", 90) - ); - - document.body.append(root); - - await vi.waitFor(() => { - expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(2); - }); - }); - - it("moves and removes exact positions through thumbnail actions", async () => { - const fixture = configured(["same.png", "middle.png", "same.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(3); - }); - const root = document.createElement("section"); - root.dataset.nodeId = "7"; - root.append( - nativePreviewImage("same.png", 0), - nativePreviewImage("middle.png", 90), - nativePreviewImage("same.png", 180) - ); - document.body.append(root); - await vi.waitFor(() => { - expect(document.querySelectorAll("[data-ss-media-move-earlier]")).toHaveLength(3); - }); - document - .querySelectorAll("[data-ss-media-move-earlier]")[2] - ?.click(); - - expect(fixture.imageWidget.value).toEqual([ - "same.png", - "same.png", - "middle.png" - ]); - expect(visiblePreviewFiles(root)).toEqual([ - "same.png", - "same.png", - "middle.png" - ]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("same.png", "same.png", "middle.png") - ); - }); - expect(fixture.onWidgetChanged).toHaveBeenCalledWith( - "image", - ["same.png", "same.png", "middle.png"], - ["same.png", "middle.png", "same.png"], - fixture.imageWidget - ); - - document - .querySelectorAll("[data-ss-media-remove]")[1] - ?.click(); - expect(fixture.imageWidget.value).toEqual(["same.png", "middle.png"]); - expect(visiblePreviewFiles(root)).toEqual(["same.png", "middle.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("same.png", "middle.png") - ); - }); - }); - - it("keeps exactly the native cells in visual order during a slow rebuild", async () => { - const delayed = deferred(); - let previewCalls = 0; - const config = { - ...CONFIG, - preview: (files: string[]) => { - previewCalls += 1; - return previewCalls === 1 - ? Promise.resolve({ images: references(...files) }) - : delayed.promise; - } - }; - const fixture = createFixture(["one.png", "two.png", "three.png"]); - configureOrderedMediaNode( - fixture.node, - fixture.app, - fixture.api, - config - ); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png", "two.png", "three.png") - ); - }); - const root = document.createElement("section"); - root.dataset.nodeId = "7"; - root.append( - nativePreviewImage("one.png", 0), - nativePreviewImage("two.png", 90), - nativePreviewImage("three.png", 180) - ); - document.body.append(root); - await vi.waitFor(() => { - expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(3); - }); - - document - .querySelectorAll("[data-ss-media-move-later]")[0] - ?.click(); - - expect(fixture.imageWidget.value).toEqual([ - "two.png", - "one.png", - "three.png" - ]); - expect(visiblePreviewFiles(root)).toEqual([ - "two.png", - "one.png", - "three.png" - ]); - expect(root.querySelectorAll("img")).toHaveLength(3); - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png", "two.png", "three.png") - ); - await new Promise((resolve) => setTimeout(resolve, 50)); - expect(visiblePreviewFiles(root)).toEqual([ - "two.png", - "one.png", - "three.png" - ]); - - delayed.resolve({ images: references("two.png", "one.png", "three.png") }); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("two.png", "one.png", "three.png") - ); - }); - expect(visiblePreviewFiles(root)).toEqual([ - "two.png", - "one.png", - "three.png" - ]); - - root.replaceChildren( - nativePreviewImage("two.png", 0), - nativePreviewImage("one.png", 90), - nativePreviewImage("three.png", 180) - ); - expect(visiblePreviewFiles(root)).toEqual([ - "two.png", - "one.png", - "three.png" - ]); - expect(root.querySelectorAll("img")).toHaveLength(3); - }); - - it("keeps the native visual and serialized order when rebuilding fails", async () => { - const logger = { warn: vi.fn() }; - let previewCalls = 0; - const config = { - ...CONFIG, - preview: (files: string[]) => { - previewCalls += 1; - return previewCalls === 1 - ? Promise.resolve({ images: references(...files) }) - : Promise.reject(new Error("preview unavailable")); - } - }; - const fixture = createFixture(["one.png", "two.png"]); - configureOrderedMediaNode( - fixture.node, - fixture.app, - fixture.api, - config, - logger - ); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png", "two.png") - ); - }); - const root = document.createElement("section"); - root.dataset.nodeId = "7"; - root.append( - nativePreviewImage("one.png", 0), - nativePreviewImage("two.png", 90) - ); - document.body.append(root); - await vi.waitFor(() => { - expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(2); - }); - - document - .querySelectorAll("[data-ss-media-move-later]")[0] - ?.click(); - - await vi.waitFor(() => { - expect(logger.warn).toHaveBeenCalledWith( - expect.stringContaining("preview unavailable"), - expect.any(Error) - ); - }); - expect(fixture.imageWidget.value).toEqual(["two.png", "one.png"]); - expect(visiblePreviewFiles(root)).toEqual(["two.png", "one.png"]); - expect(root.querySelectorAll("img")).toHaveLength(2); - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png", "two.png") - ); - }); - - it("preserves native thumbnail inspection clicks in Nodes 2.0", async () => { - const fixture = configured(["one.png", "two.png", "three.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(3); - }); - const root = document.createElement("section"); - root.dataset.nodeId = "7"; - const images = ["one.png", "two.png", "three.png"].map( - (filename, index) => nativePreviewImage(filename, index * 90) - ); - root.append(...images); - document.body.append(root); - const clicked = vi.fn(); - images[1]?.addEventListener("click", clicked); - - images[1]?.dispatchEvent(new MouseEvent("click", { bubbles: true })); - expect(clicked).toHaveBeenCalledOnce(); - }); - - it("moves Nodes 1.0 native canvas cells through imageRects", async () => { - const fixture = createFixture(["one.png", "two.png"]); - const canvas = document.createElement("canvas"); - mockRect(canvas, 0, 0, 500, 500); - document.body.append(canvas); - fixture.app.canvas = { - canvas, - convertEventToCanvasOffset: (event) => [event.clientX, event.clientY], - convertOffsetToCanvas: (position) => [...position] - }; - Object.assign(fixture.node, { - pos: [10, 100], - imageIndex: null, - imageRects: [ - [0, 50, 80, 80], - [80, 50, 80, 80] - ], - imgs: [nativePreviewImage("one.png", 0), nativePreviewImage("two.png", 90)] - }); - configureOrderedMediaNode(fixture.node, fixture.app, fixture.api, CONFIG); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toHaveLength(2); - }); - await vi.waitFor(() => { - expect(document.querySelectorAll("[data-ss-media-move-later]")).toHaveLength(2); - }); - document - .querySelectorAll("[data-ss-media-move-later]")[0] - ?.click(); - - expect(fixture.imageWidget.value).toEqual(["two.png", "one.png"]); - expect(imageFiles(fixture.node.imgs)).toEqual(["two.png", "one.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("two.png", "one.png") - ); - }); - }); - - it("normalizes scalar workflows and survives Nodes 2.0 reactive assignments", async () => { - const fixture = configured("saved.png", true); - - const onGraphConfigured = fixture.node.onGraphConfigured; - if (!onGraphConfigured) throw new Error("Expected graph configuration callback."); - onGraphConfigured(); - - expect(fixture.imageWidget.value).toEqual(["saved.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("saved.png") - ); - }); - }); - - it("retains array-valued media when the legacy graph lifecycle clears the widget", async () => { - const fixture = createFixture(["one.png", "two.png"]); - fixture.node.onGraphConfigured = () => { - fixture.imageWidget.value = undefined; - }; - configureOrderedMediaNode( - fixture.node, - fixture.app, - fixture.api, - CONFIG - ); - fixture.node.onGraphConfigured(); - - expect(fixture.imageWidget.value).toEqual(["one.png", "two.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png", "two.png") - ); - }); - }); - - it("restores persisted widget values when Comfy clears the live media widget", async () => { - const fixture = createFixture(undefined); - fixture.node.widgets_values = [ - ["one.png", "two.png"], - "image", - "ordered_media", - "ordered_media" - ]; - configureOrderedMediaNode( - fixture.node, - fixture.app, - fixture.api, - CONFIG - ); - fixture.imageWidget.value = undefined; - - fixture.node.onGraphConfigured?.(); - - expect(fixture.imageWidget.value).toEqual(["one.png", "two.png"]); - await vi.waitFor(() => { - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual( - references("one.png", "two.png") - ); - }); - }); - - it("restores native handlers and clears preview on node removal", () => { - const fixture = createFixture(["one.png"]); - const paste = fixture.node.pasteFiles; - const drop = fixture.node.onDragDrop; - configureOrderedMediaNode( - fixture.node, - fixture.app, - fixture.api, - CONFIG - ); - - fixture.node.onRemoved(); - - expect(fixture.node.pasteFiles).toBe(paste); - expect(fixture.node.onDragDrop).toBe(drop); - expect(fixture.originalRemoved).toHaveBeenCalledOnce(); - expect(fixture.app.nodeOutputs?.["7"]?.images).toEqual([]); - }); - - it("registers a guarded extension and ignores unrelated nodes", async () => { - const fixture = createFixture([]); - let extension: ComfyExtension | undefined; - fixture.app.registerExtension = (value: ComfyExtension) => { - extension = value; - }; - registerOrderedMediaNode( - fixture.app, - fixture.api, - "test.extension", - CONFIG - ); - fixture.node.constructor.comfyClass = "Other.Node"; - - await extension?.nodeCreated?.(fixture.node); - - expect(extension?.name).toBe("test.extension"); - expect(fixture.node.addWidget).not.toHaveBeenCalled(); - }); -}); - -interface TestWidget { - name: string; - value: unknown; - type?: string; - label?: string; - hidden?: boolean; - disabled?: boolean; - callback?: (value?: unknown) => void; - computeSize?: (width?: number) => [number, number]; - options?: { - canvasOnly?: boolean; - serialize?: boolean; - tooltip?: string; - hidden?: boolean; - }; -} - -function configured(value: unknown, reactive = false) { - const fixture = createFixture(value, reactive); - configureOrderedMediaNode( - fixture.node, - fixture.app, - fixture.api, - CONFIG - ); - return fixture; -} - -function createFixture(value: unknown, reactive = false) { - const nativeUpload = vi.fn(); - const imageWidget: TestWidget = { - name: "image", - value, - type: "combo", - options: { canvasOnly: true } - }; - if (reactive) { - let stored = value; - Object.defineProperty(imageWidget, "value", { - configurable: true, - get: () => stored, - set: (next: unknown) => { - stored = next; - imageWidget.callback?.(next); - } - }); - } - const uploadWidget: TestWidget = { - name: "upload", - value: "image", - type: "button", - callback: nativeUpload, - options: { serialize: false, canvasOnly: true } - }; - const widgets = [imageWidget, uploadWidget]; - const originalRemoved = vi.fn(); - const onWidgetChanged = vi.fn(); - const node = { - constructor: { comfyClass: "SimpleSyrup.LoadImageList" }, - id: 7, - widgets, - widgets_values: undefined as unknown[] | undefined, - imgs: undefined as HTMLImageElement[] | undefined, - graph: { setDirtyCanvas: vi.fn() }, - pasteFiles: vi.fn(() => true) as (...args: unknown[]) => unknown, - onDragDrop: vi.fn(() => true) as (...args: unknown[]) => unknown, - onRemoved: originalRemoved as (...args: unknown[]) => unknown, - onGraphConfigured: undefined as ((...args: unknown[]) => unknown) | undefined, - onWidgetChanged, - addWidget: vi.fn( - ( - type: "button" | "combo", - name: string, - widgetValue: string, - callback: (value?: unknown) => void, - options?: TestWidget["options"] - ): TestWidget => { - const created: TestWidget = { - type, - name, - value: widgetValue, - callback - }; - if (options) created.options = options; - widgets.push(created); - return created; - } - ) - }; - activeNodes.push(node); - const app = { - nodeOutputs: {}, - ui: { settings: { addSetting: vi.fn() } }, - registerExtension: vi.fn() - } as unknown as ComfyApp; - const events = new EventTarget(); - const api = events as ComfyApi; - const executed = vi.fn(); - events.addEventListener("executed", (event: Event) => { - executed((event as CustomEvent).detail); - }); - return { - api, - app, - executed, - imageWidget, - nativeUpload, - node, - onWidgetChanged, - originalRemoved, - uploadWidget - }; -} - -function nativePreviewImage( - filename: string, - left: number -): HTMLImageElement { - const image = document.createElement("img"); - image.src = `/api/view?filename=${filename}&subfolder=&type=input&preview=webp`; - Object.defineProperties(image, { - complete: { configurable: true, value: true }, - naturalWidth: { configurable: true, value: 80 } - }); - mockRect(image, left, 0, 80, 80); - return image; -} - -function nativePreviewButton(filename: string, left: number): HTMLButtonElement { - const button = document.createElement("button"); - button.append(nativePreviewImage(filename, left)); - mockRect(button, left, 0, 80, 80); - return button; -} - -function visiblePreviewFiles(root: HTMLElement): string[] { - return Array.from(root.querySelectorAll("img")) - .filter((image) => image.style.display !== "none") - .map((image) => new URL(image.src).searchParams.get("filename") ?? ""); -} - -function imageFiles(images: HTMLImageElement[] | undefined): string[] { - return (images ?? []).map( - (image) => new URL(image.src).searchParams.get("filename") ?? "" - ); -} - -function mockRect( - element: Element, - left: number, - top: number, - width: number, - height: number -): void { - element.getBoundingClientRect = () => new DOMRect(left, top, width, height); -} - -function widget( - fixture: ReturnType, - name: string -): TestWidget { - const found = fixture.node.widgets.find((candidate) => candidate.name === name); - if (!found) throw new Error(`Missing test widget ${name}.`); - return found; -} - -function selectFiles( - fixture: ReturnType, - files: string[] -): void { - fixture.imageWidget.value = files; - fixture.imageWidget.callback?.(files); -} - -function references(...filenames: string[]): ComfyImageResult[] { - return filenames.map((filename) => ({ - filename, - subfolder: "", - type: "input" - })); -} - -function deferred(): { - readonly promise: Promise; - readonly resolve: (value: T) => void; -} { - let resolvePromise: ((value: T) => void) | undefined; - const promise = new Promise((resolve) => { - resolvePromise = resolve; - }); - return { - promise, - resolve: (value) => { - if (!resolvePromise) throw new Error("Deferred promise is unavailable."); - resolvePromise(value); - } - }; -} diff --git a/web/tests/maskAtlas.test.ts b/web/tests/segmentation/maskAtlas.test.ts similarity index 94% rename from web/tests/maskAtlas.test.ts rename to web/tests/segmentation/maskAtlas.test.ts index ebd77bb..fc44683 100644 --- a/web/tests/maskAtlas.test.ts +++ b/web/tests/segmentation/maskAtlas.test.ts @@ -4,8 +4,8 @@ import { describe, expect, it } from "vitest"; -import { MaskAtlas } from "../src/maskAtlas"; -import type { SegPreviewDocument } from "../src/segPreviewTypes"; +import { MaskAtlas } from "../../src/maskAtlas"; +import type { SegPreviewDocument } from "../../src/segPreviewTypes"; describe("MaskAtlas", () => { it("returns nested hits from smallest to largest and respects mask holes", () => { diff --git a/web/tests/maskBatchUpload.test.ts b/web/tests/segmentation/maskBatchUpload.test.ts similarity index 98% rename from web/tests/maskBatchUpload.test.ts rename to web/tests/segmentation/maskBatchUpload.test.ts index bbfc2a4..8c24043 100644 --- a/web/tests/maskBatchUpload.test.ts +++ b/web/tests/segmentation/maskBatchUpload.test.ts @@ -7,8 +7,8 @@ import { afterEach, describe, expect, it, vi } from "vitest"; import { configureMaskBatchNode, registerMaskBatchUpload -} from "../src/maskBatchUpload"; -import type { ComfyApi, ComfyApp, ComfyExtension } from "../src/types"; +} from "../../src/maskBatchUpload"; +import type { ComfyApi, ComfyApp, ComfyExtension } from "../../src/types"; afterEach(() => { document.body.replaceChildren(); diff --git a/web/tests/maskedTextureRenderer.test.ts b/web/tests/segmentation/maskedTextureRenderer.test.ts similarity index 95% rename from web/tests/maskedTextureRenderer.test.ts rename to web/tests/segmentation/maskedTextureRenderer.test.ts index 9f11d00..215be49 100644 --- a/web/tests/maskedTextureRenderer.test.ts +++ b/web/tests/segmentation/maskedTextureRenderer.test.ts @@ -4,7 +4,7 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { MaskedTextureRenderer } from "../src/maskedTextureRenderer"; +import { MaskedTextureRenderer } from "../../src/maskedTextureRenderer"; describe("MaskedTextureRenderer", () => { afterEach(() => { diff --git a/web/tests/segPreviewInspector.test.ts b/web/tests/segmentation/segPreviewInspector.test.ts similarity index 95% rename from web/tests/segPreviewInspector.test.ts rename to web/tests/segmentation/segPreviewInspector.test.ts index 7ecccb7..249d9ea 100644 --- a/web/tests/segPreviewInspector.test.ts +++ b/web/tests/segmentation/segPreviewInspector.test.ts @@ -4,9 +4,9 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { SegPreviewInspector } from "../src/segPreviewInspector"; -import type { SegPreviewDocument } from "../src/segPreviewTypes"; -import { installCanvasMock } from "./testUtils"; +import { SegPreviewInspector } from "../../src/segPreviewInspector"; +import type { SegPreviewDocument } from "../../src/segPreviewTypes"; +import { installCanvasMock } from "../support/testUtils"; describe("SegPreviewInspector", () => { afterEach(() => { diff --git a/web/tests/segPreviewNativeSurface.test.ts b/web/tests/segmentation/segPreviewNativeSurface.test.ts similarity index 96% rename from web/tests/segPreviewNativeSurface.test.ts rename to web/tests/segmentation/segPreviewNativeSurface.test.ts index c0fd50f..83cf2f6 100644 --- a/web/tests/segPreviewNativeSurface.test.ts +++ b/web/tests/segmentation/segPreviewNativeSurface.test.ts @@ -4,9 +4,9 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { SegPreviewNativeSurface } from "../src/segPreviewNativeSurface"; -import type { ComfyApi } from "../src/types"; -import { createFakeComfyApp } from "./testUtils"; +import { SegPreviewNativeSurface } from "../../src/segPreviewNativeSurface"; +import type { ComfyApi } from "../../src/types"; +import { createFakeComfyApp } from "../support/testUtils"; describe("SegPreviewNativeSurface", () => { afterEach(() => { diff --git a/web/tests/segPreviewNode.test.ts b/web/tests/segmentation/segPreviewNode.test.ts similarity index 97% rename from web/tests/segPreviewNode.test.ts rename to web/tests/segmentation/segPreviewNode.test.ts index 29f9686..27438aa 100644 --- a/web/tests/segPreviewNode.test.ts +++ b/web/tests/segmentation/segPreviewNode.test.ts @@ -4,13 +4,13 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { registerSimplePreviewSEGS } from "../src/segPreviewNode"; +import { registerSimplePreviewSEGS } from "../../src/segPreviewNode"; import type { ComfyApi, ComfyNodeExecutionOutput, ComfyExtension -} from "../src/types"; -import { createFakeComfyApp, installCanvasMock } from "./testUtils"; +} from "../../src/types"; +import { createFakeComfyApp, installCanvasMock } from "../support/testUtils"; describe("Simple Preview SEGS node integration", () => { afterEach(() => { diff --git a/web/tests/segPreviewTypes.test.ts b/web/tests/segmentation/segPreviewTypes.test.ts similarity index 92% rename from web/tests/segPreviewTypes.test.ts rename to web/tests/segmentation/segPreviewTypes.test.ts index 89a4f51..b0286f1 100644 --- a/web/tests/segPreviewTypes.test.ts +++ b/web/tests/segmentation/segPreviewTypes.test.ts @@ -4,8 +4,8 @@ import { describe, expect, it } from "vitest"; -import { comfyImageUrl } from "../src/comfyImageUrl"; -import { parseSegPreviewDocument } from "../src/segPreviewTypes"; +import { comfyImageUrl } from "../../src/comfyImageUrl"; +import { parseSegPreviewDocument } from "../../src/segPreviewTypes"; describe("SEG preview transport", () => { it("validates execution payloads and uses the latest mapped document", () => { diff --git a/web/tests/settings.test.ts b/web/tests/settings.test.ts deleted file mode 100644 index 5cb1684..0000000 --- a/web/tests/settings.test.ts +++ /dev/null @@ -1,510 +0,0 @@ -// SimpleSyrup - workflow-focused ComfyUI extensions for image generation -// Copyright (C) 2026 Artificial Sweetener and contributors -// SPDX-License-Identifier: AGPL-3.0-or-later - -import { describe, expect, it, vi } from "vitest"; - -import { - SIMPLE_SYRUP_SETTING_ID, - SIMPLE_SYRUP_SETTING_DESCRIPTION, - SIMPLE_SYRUP_SETTING_LABEL -} from "../src/downloadableModelsSetting"; -import { - EXTERNAL_LLM_API_KEY_SETTING_ID, - EXTERNAL_LLM_ENDPOINT_SETTING_ID -} from "../src/externalLlmSettings"; -import { - registerSimpleSyrupSettings -} from "../src/settingsRegistration"; -import type { SimpleSyrupSettingsApi } from "../src/settingsRegistration"; -import { QUANT_CACHE_SETTING_ID } from "../src/quantCacheSetting"; -import { createFakeComfyApp } from "./testUtils"; - -describe("Comfy settings registration", () => { - it("registers the ShowDownloadableModels setting from backend state", async () => { - const app = createFakeComfyApp(); - const api = fakeSettingsApi(false); - - await registerSimpleSyrupSettings(app, api); - - expect(app.ui.settings.definitions).toHaveLength(4); - expect(app.ui.settings.definitions[0]).toMatchObject({ - id: SIMPLE_SYRUP_SETTING_ID, - name: SIMPLE_SYRUP_SETTING_LABEL, - type: "boolean", - defaultValue: false, - tooltip: SIMPLE_SYRUP_SETTING_DESCRIPTION - }); - expect(SIMPLE_SYRUP_SETTING_DESCRIPTION).toContain("WD14 tagger"); - expect(SIMPLE_SYRUP_SETTING_DESCRIPTION).toContain("Ultralytics"); - expect(app.ui.settings.settings[0]?.value).toBe(false); - expect(app.ui.settings.definitions[1]).toMatchObject({ - id: QUANT_CACHE_SETTING_ID, - sortOrder: 321 - }); - expect(typeof app.ui.settings.definitions[1]?.type).toBe("function"); - expect(app.ui.settings.definitions[2]).toMatchObject({ - id: EXTERNAL_LLM_ENDPOINT_SETTING_ID, - sortOrder: 320 - }); - expect(typeof app.ui.settings.definitions[2]?.type).toBe("function"); - expect(app.ui.settings.definitions[3]).toMatchObject({ - id: EXTERNAL_LLM_API_KEY_SETTING_ID, - sortOrder: 319 - }); - expect(typeof app.ui.settings.definitions[3]?.type).toBe("function"); - }); - - it("saves setting changes to the backend", async () => { - const app = createFakeComfyApp(); - const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); - app.refreshComboInNodes = refreshComboInNodes; - const saveSettings = vi - .fn() - .mockResolvedValue({ - show_downloadable_models: true, - quant_cache_limit_gib: 20 - }); - const api: SimpleSyrupSettingsApi = { - getSettings: vi.fn().mockResolvedValue({ - show_downloadable_models: false, - quant_cache_limit_gib: 20 - }), - saveSettings, - getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), - saveExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), - saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) - }; - - await registerSimpleSyrupSettings(app, api); - await app.ui.settings.definitions[0]?.onChange?.(true); - - expect(saveSettings).toHaveBeenCalledWith({ - show_downloadable_models: true, - quant_cache_limit_gib: 20 - }); - expect(app.ui.settings.settings[0]?.value).toBe(true); - expect(refreshComboInNodes).toHaveBeenCalledOnce(); - }); - - it("falls back to the default and warns when backend load fails", async () => { - const app = createFakeComfyApp(); - const logger = { warn: vi.fn() }; - const api: SimpleSyrupSettingsApi = { - getSettings: vi.fn().mockRejectedValue(new Error("offline")), - saveSettings: vi.fn().mockResolvedValue({ - show_downloadable_models: true, - quant_cache_limit_gib: 20 - }), - getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), - saveExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), - saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) - }; - - await registerSimpleSyrupSettings(app, api, logger); - - expect(app.ui.settings.settings[0]?.value).toBe(true); - expect(logger.warn).toHaveBeenCalledWith( - expect.stringContaining("Could not load SimpleSyrup settings"), - expect.any(Error) - ); - }); - - it("warns and restores the previous value when backend save fails", async () => { - const app = createFakeComfyApp(); - const logger = { warn: vi.fn() }; - const saveSettings = vi - .fn() - .mockResolvedValueOnce({ - show_downloadable_models: true, - quant_cache_limit_gib: 20 - }) - .mockRejectedValueOnce(new Error("rejected")); - const api: SimpleSyrupSettingsApi = { - getSettings: vi.fn().mockResolvedValue({ - show_downloadable_models: false, - quant_cache_limit_gib: 20 - }), - saveSettings, - getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), - saveExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), - saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) - }; - - await registerSimpleSyrupSettings(app, api, logger); - await app.ui.settings.definitions[0]?.onChange?.(true); - await app.ui.settings.definitions[0]?.onChange?.(false); - - expect(logger.warn).toHaveBeenCalledWith( - expect.stringContaining("Could not save SimpleSyrup settings"), - expect.any(Error) - ); - expect(app.ui.settings.settings[0]?.value).toBe(true); - }); - - it("keeps a saved setting when live model-choice refresh fails", async () => { - const app = createFakeComfyApp(); - const logger = { warn: vi.fn() }; - app.refreshComboInNodes = vi.fn().mockRejectedValue(new Error("offline")); - const api = fakeSettingsApi(false); - - await registerSimpleSyrupSettings(app, api, logger); - await app.ui.settings.definitions[0]?.onChange?.(true); - - expect(app.ui.settings.settings[0]?.value).toBe(true); - expect(logger.warn).toHaveBeenCalledWith( - expect.stringContaining("Could not refresh Comfy loader model choices"), - expect.any(Error) - ); - }); - - it("shows global quant cache usage and saves its GiB limit", async () => { - const app = createFakeComfyApp(); - const api = fakeSettingsApi(true); - const saveSettings = vi - .fn() - .mockResolvedValue({ - show_downloadable_models: true, - quant_cache_limit_gib: 30 - }); - api.saveSettings = saveSettings; - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 1)); - const input = requiredInput(control, "input[type=number]"); - const saveButton = requiredButton(control, "button"); - - expect(input.value).toBe("20"); - expect(control.textContent).toContain("models/SyrupQuants"); - input.value = "30"; - saveButton.click(); - await flushPromises(); - - expect(saveSettings).toHaveBeenCalledWith({ - show_downloadable_models: true, - quant_cache_limit_gib: 30 - }); - expect(input.value).toBe("30"); - }); - - it("clears inactive quant artifacts and refreshes cache status", async () => { - const app = createFakeComfyApp(); - const api = fakeSettingsApi(true); - const clearQuantCache = vi.fn().mockResolvedValue({ - ...defaultQuantCacheStatus(), - removed_artifacts: 2, - removed_bytes: 1024 - }); - api.clearQuantCache = clearQuantCache; - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 1)); - const buttons = control.querySelectorAll("button"); - const clearButton = buttons[1]; - if (!(clearButton instanceof HTMLButtonElement)) { - throw new Error("Expected quant cache clear button."); - } - clearButton.click(); - await flushPromises(); - - expect(clearQuantCache).toHaveBeenCalledOnce(); - expect(control.textContent).toContain("0.00 GiB used"); - }); - - it("restores the previous quant limit when backend saving fails", async () => { - const app = createFakeComfyApp(); - const logger = { warn: vi.fn() }; - const api = fakeSettingsApi(true); - api.saveSettings = vi.fn().mockRejectedValue(new Error("rejected")); - - await registerSimpleSyrupSettings(app, api, logger); - const control = renderSetting(getDefinition(app, 1)); - const input = requiredInput(control, "input[type=number]"); - const saveButton = requiredButton(control, "button"); - input.value = "30"; - saveButton.click(); - await flushPromises(); - - expect(input.value).toBe("20"); - expect(control.textContent).toContain("was not saved"); - expect(logger.warn).toHaveBeenCalledWith( - expect.stringContaining("quant cache limit"), - expect.any(Error) - ); - }); - - it("saves external LLM endpoint changes to the backend", async () => { - const app = createFakeComfyApp(); - const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); - app.refreshComboInNodes = refreshComboInNodes; - const saveExternalLLMSettings = - vi.fn().mockResolvedValue({ - ...defaultExternalLLMSettings(), - base_url: "https://provider.example/v1" - }); - const api = fakeSettingsApi(true); - api.saveExternalLLMSettings = saveExternalLLMSettings; - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 2)); - const input = requiredInput(control, "input"); - const button = requiredButton(control, "button"); - - input.value = "https://provider.example/v1"; - button.click(); - await flushPromises(); - - expect(saveExternalLLMSettings).toHaveBeenCalledWith({ - base_url: "https://provider.example/v1", - default_model: "" - }); - expect(refreshComboInNodes).toHaveBeenCalledOnce(); - expect(control.textContent).toContain("Endpoint saved."); - }); - - it("keeps partial endpoint text visible while the user is typing", async () => { - const app = createFakeComfyApp(); - const saveExternalLLMSettings = - vi.fn(); - const api = fakeSettingsApi(true); - api.saveExternalLLMSettings = saveExternalLLMSettings; - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 2)); - const input = requiredInput(control, "input"); - const button = requiredButton(control, "button"); - - input.value = "https://"; - button.click(); - await Promise.resolve(); - - expect(saveExternalLLMSettings).not.toHaveBeenCalled(); - expect(input.value).toBe("https://"); - }); - - it("saves API keys through an add dialog without showing the stored key", async () => { - const app = createFakeComfyApp(); - const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); - app.refreshComboInNodes = refreshComboInNodes; - const saveExternalLLMApiKey = - vi.fn().mockResolvedValue({ - ...configuredExternalLLMSettings(), - has_api_key: true - }); - const api = fakeSettingsApi(true); - api.getExternalLLMSettings = vi - .fn() - .mockResolvedValue(configuredExternalLLMSettings()); - api.saveExternalLLMApiKey = saveExternalLLMApiKey; - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 3)); - const addButton = requiredButton(control, "button"); - expect(addButton.textContent).toBe("Add API Key"); - - addButton.click(); - const dialogInput = requiredInput(document.body, ".simple-syrup-dialog input"); - const submitButton = requiredButton( - document.body, - ".simple-syrup-dialog button" - ); - - dialogInput.value = "secret"; - submitButton.click(); - await flushPromises(); - - expect(saveExternalLLMApiKey).toHaveBeenCalledWith({ api_key: "secret" }); - expect(refreshComboInNodes).toHaveBeenCalledOnce(); - expect(document.body.querySelector(".simple-syrup-dialog")).toBeNull(); - expect(control.textContent).toContain("API key remembered."); - expect(control.textContent).not.toContain("secret"); - }); - - it("keeps save success visible if model choice refresh fails", async () => { - const app = createFakeComfyApp(); - const logger = { warn: vi.fn() }; - app.refreshComboInNodes = vi.fn().mockRejectedValue(new Error("offline")); - const saveExternalLLMSettings = - vi.fn().mockResolvedValue({ - ...defaultExternalLLMSettings(), - base_url: "https://provider.example/v1" - }); - const api = fakeSettingsApi(true); - api.saveExternalLLMSettings = saveExternalLLMSettings; - - await registerSimpleSyrupSettings(app, api, logger); - const control = renderSetting(getDefinition(app, 2)); - const input = requiredInput(control, "input"); - const button = requiredButton(control, "button"); - - input.value = "https://provider.example/v1"; - button.click(); - await flushPromises(); - - expect(control.textContent).toContain("Endpoint saved."); - expect(logger.warn).toHaveBeenCalledWith( - expect.stringContaining("Could not refresh Comfy node definitions"), - expect.any(Error) - ); - }); - - it("requires a saved endpoint before opening the API key dialog", async () => { - const app = createFakeComfyApp(); - const api = fakeSettingsApi(true); - const saveExternalLLMApiKey = vi.fn< - SimpleSyrupSettingsApi["saveExternalLLMApiKey"] - >(); - api.saveExternalLLMApiKey = saveExternalLLMApiKey; - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 3)); - const addButton = requiredButton(control, "button"); - - addButton.click(); - - expect(document.body.querySelector(".simple-syrup-dialog")).toBeNull(); - expect(control.textContent).toContain("Save endpoint first."); - expect(saveExternalLLMApiKey).not.toHaveBeenCalled(); - }); - - it("shows backend API key save errors in the setting row", async () => { - const app = createFakeComfyApp(); - const api = fakeSettingsApi(true); - api.getExternalLLMSettings = vi - .fn() - .mockResolvedValue(configuredExternalLLMSettings()); - api.saveExternalLLMApiKey = vi - .fn() - .mockRejectedValue(new Error("Configure an external LLM endpoint first.")); - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 3)); - const addButton = requiredButton(control, "button"); - - addButton.click(); - const dialogInput = requiredInput(document.body, ".simple-syrup-dialog input"); - const submitButton = requiredButton( - document.body, - ".simple-syrup-dialog button" - ); - - dialogInput.value = "secret"; - submitButton.click(); - await Promise.resolve(); - await Promise.resolve(); - - expect(control.textContent).toContain("Save endpoint first."); - expect(control.textContent).not.toContain("secret"); - }); - - it("offers replacement when an API key is remembered", async () => { - const app = createFakeComfyApp(); - const api = fakeSettingsApi(true); - api.getExternalLLMSettings = vi - .fn() - .mockResolvedValue({ ...defaultExternalLLMSettings(), has_api_key: true }); - - await registerSimpleSyrupSettings(app, api); - const control = renderSetting(getDefinition(app, 3)); - const button = control.querySelector("button"); - - expect(button?.textContent).toBe("Replace API Key"); - expect(control.textContent).toContain("API key remembered."); - }); -}); - -function fakeSettingsApi( - showDownloadableModels: boolean -): SimpleSyrupSettingsApi { - return { - getSettings: vi.fn().mockResolvedValue({ - show_downloadable_models: showDownloadableModels, - quant_cache_limit_gib: 20 - }), - saveSettings: vi.fn().mockImplementation((settings) => Promise.resolve(settings)), - getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), - getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), - saveExternalLLMSettings: vi - .fn() - .mockImplementation((settings) => - Promise.resolve({ ...defaultExternalLLMSettings(), ...settings }) - ), - saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) - }; -} - -function defaultQuantCacheStatus() { - return { - path: "models/SyrupQuants", - usage_bytes: 0, - limit_bytes: 20 * 1024 ** 3, - artifact_count: 0, - active_artifact_count: 0 - }; -} - -function defaultExternalLLMSettings() { - return { - base_url: "", - cached_models: [], - default_model: "", - has_api_key: false - }; -} - -function configuredExternalLLMSettings() { - return { - base_url: "https://provider.example/v1", - cached_models: [], - default_model: "", - has_api_key: false - }; -} - -function renderSetting(definition: { - type: "boolean" | "text" | (() => HTMLElement); -}): HTMLElement { - expect(typeof definition.type).toBe("function"); - return (definition.type as () => HTMLElement)(); -} - -function getDefinition(app: ReturnType, index: number) { - const definition = app.ui.settings.definitions[index]; - if (!definition) { - throw new Error(`Expected setting definition at index ${String(index)}.`); - } - return definition; -} - -function requiredInput(parent: ParentNode, selector: string): HTMLInputElement { - const input = parent.querySelector(selector); - if (!(input instanceof HTMLInputElement)) { - throw new Error(`Expected input for selector ${selector}.`); - } - return input; -} - -function requiredButton(parent: ParentNode, selector: string): HTMLButtonElement { - const button = parent.querySelector(selector); - if (!(button instanceof HTMLButtonElement)) { - throw new Error(`Expected button for selector ${selector}.`); - } - return button; -} - -async function flushPromises(): Promise { - for (let index = 0; index < 6; index += 1) { - await Promise.resolve(); - } -} diff --git a/web/tests/settings/externalLlmSettings.test.ts b/web/tests/settings/externalLlmSettings.test.ts new file mode 100644 index 0000000..be7fa20 --- /dev/null +++ b/web/tests/settings/externalLlmSettings.test.ts @@ -0,0 +1,188 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { describe, expect, it, vi } from "vitest"; + +import { registerSimpleSyrupSettings } from "../../src/settingsRegistration"; +import type { SimpleSyrupSettingsApi } from "../../src/settingsRegistration"; +import { createFakeComfyApp } from "../support/testUtils"; +import { + configuredExternalLLMSettings, + defaultExternalLLMSettings, + fakeSettingsApi, + flushPromises, + getDefinition, + renderSetting, + requiredButton, + requiredInput +} from "./settingsTestSupport"; + +describe("external LLM settings registration", () => { + it("saves endpoint changes to the backend", async () => { + const app = createFakeComfyApp(); + const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); + app.refreshComboInNodes = refreshComboInNodes; + const saveExternalLLMSettings = + vi.fn().mockResolvedValue({ + ...defaultExternalLLMSettings(), + base_url: "https://provider.example/v1" + }); + const api = fakeSettingsApi(true); + api.saveExternalLLMSettings = saveExternalLLMSettings; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 2)); + const input = requiredInput(control, "input"); + const button = requiredButton(control, "button"); + input.value = "https://provider.example/v1"; + button.click(); + await flushPromises(); + + expect(saveExternalLLMSettings).toHaveBeenCalledWith({ + base_url: "https://provider.example/v1", + default_model: "" + }); + expect(refreshComboInNodes).toHaveBeenCalledOnce(); + expect(control.textContent).toContain("Endpoint saved."); + }); + + it("keeps partial endpoint text visible while the user is typing", async () => { + const app = createFakeComfyApp(); + const saveExternalLLMSettings = + vi.fn(); + const api = fakeSettingsApi(true); + api.saveExternalLLMSettings = saveExternalLLMSettings; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 2)); + const input = requiredInput(control, "input"); + const button = requiredButton(control, "button"); + input.value = "https://"; + button.click(); + await Promise.resolve(); + + expect(saveExternalLLMSettings).not.toHaveBeenCalled(); + expect(input.value).toBe("https://"); + }); + + it("saves API keys without showing the stored key", async () => { + const app = createFakeComfyApp(); + const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); + app.refreshComboInNodes = refreshComboInNodes; + const saveExternalLLMApiKey = + vi.fn().mockResolvedValue({ + ...configuredExternalLLMSettings(), + has_api_key: true + }); + const api = fakeSettingsApi(true); + api.getExternalLLMSettings = vi + .fn() + .mockResolvedValue(configuredExternalLLMSettings()); + api.saveExternalLLMApiKey = saveExternalLLMApiKey; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 3)); + const addButton = requiredButton(control, "button"); + expect(addButton.textContent).toBe("Add API Key"); + addButton.click(); + const dialogInput = requiredInput(document.body, ".simple-syrup-dialog input"); + const submitButton = requiredButton( + document.body, + ".simple-syrup-dialog button" + ); + dialogInput.value = "secret"; + submitButton.click(); + await flushPromises(); + + expect(saveExternalLLMApiKey).toHaveBeenCalledWith({ api_key: "secret" }); + expect(refreshComboInNodes).toHaveBeenCalledOnce(); + expect(document.body.querySelector(".simple-syrup-dialog")).toBeNull(); + expect(control.textContent).toContain("API key remembered."); + expect(control.textContent).not.toContain("secret"); + }); + + it("keeps save success visible if model choice refresh fails", async () => { + const app = createFakeComfyApp(); + const logger = { warn: vi.fn() }; + app.refreshComboInNodes = vi.fn().mockRejectedValue(new Error("offline")); + const saveExternalLLMSettings = + vi.fn().mockResolvedValue({ + ...defaultExternalLLMSettings(), + base_url: "https://provider.example/v1" + }); + const api = fakeSettingsApi(true); + api.saveExternalLLMSettings = saveExternalLLMSettings; + + await registerSimpleSyrupSettings(app, api, logger); + const control = renderSetting(getDefinition(app, 2)); + const input = requiredInput(control, "input"); + const button = requiredButton(control, "button"); + input.value = "https://provider.example/v1"; + button.click(); + await flushPromises(); + + expect(control.textContent).toContain("Endpoint saved."); + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("Could not refresh Comfy node definitions"), + expect.any(Error) + ); + }); + + it("requires a saved endpoint before opening the API key dialog", async () => { + const app = createFakeComfyApp(); + const api = fakeSettingsApi(true); + const saveExternalLLMApiKey = vi.fn< + SimpleSyrupSettingsApi["saveExternalLLMApiKey"] + >(); + api.saveExternalLLMApiKey = saveExternalLLMApiKey; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 3)); + requiredButton(control, "button").click(); + + expect(document.body.querySelector(".simple-syrup-dialog")).toBeNull(); + expect(control.textContent).toContain("Save endpoint first."); + expect(saveExternalLLMApiKey).not.toHaveBeenCalled(); + }); + + it("shows backend API key save errors in the setting row", async () => { + const app = createFakeComfyApp(); + const api = fakeSettingsApi(true); + api.getExternalLLMSettings = vi + .fn() + .mockResolvedValue(configuredExternalLLMSettings()); + api.saveExternalLLMApiKey = vi + .fn() + .mockRejectedValue(new Error("Configure an external LLM endpoint first.")); + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 3)); + requiredButton(control, "button").click(); + const dialogInput = requiredInput(document.body, ".simple-syrup-dialog input"); + const submitButton = requiredButton( + document.body, + ".simple-syrup-dialog button" + ); + dialogInput.value = "secret"; + submitButton.click(); + await Promise.resolve(); + await Promise.resolve(); + + expect(control.textContent).toContain("Save endpoint first."); + expect(control.textContent).not.toContain("secret"); + }); + + it("offers replacement when an API key is remembered", async () => { + const app = createFakeComfyApp(); + const api = fakeSettingsApi(true); + api.getExternalLLMSettings = vi + .fn() + .mockResolvedValue({ ...defaultExternalLLMSettings(), has_api_key: true }); + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 3)); + expect(control.querySelector("button")?.textContent).toBe("Replace API Key"); + expect(control.textContent).toContain("API key remembered."); + }); +}); diff --git a/web/tests/settings/settings.test.ts b/web/tests/settings/settings.test.ts new file mode 100644 index 0000000..a2d8ad9 --- /dev/null +++ b/web/tests/settings/settings.test.ts @@ -0,0 +1,211 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { describe, expect, it, vi } from "vitest"; + +import { + SIMPLE_SYRUP_SETTING_DESCRIPTION, + SIMPLE_SYRUP_SETTING_ID, + SIMPLE_SYRUP_SETTING_LABEL +} from "../../src/downloadableModelsSetting"; +import { + EXTERNAL_LLM_API_KEY_SETTING_ID, + EXTERNAL_LLM_ENDPOINT_SETTING_ID +} from "../../src/externalLlmSettings"; +import { QUANT_CACHE_SETTING_ID } from "../../src/quantCacheSetting"; +import { registerSimpleSyrupSettings } from "../../src/settingsRegistration"; +import type { SimpleSyrupSettingsApi } from "../../src/settingsRegistration"; +import { createFakeComfyApp } from "../support/testUtils"; +import { + defaultExternalLLMSettings, + defaultQuantCacheStatus, + fakeSettingsApi, + flushPromises, + getDefinition, + renderSetting, + requiredButton, + requiredInput +} from "./settingsTestSupport"; + +describe("general and quant-cache settings registration", () => { + it("registers every setting from backend state", async () => { + const app = createFakeComfyApp(); + await registerSimpleSyrupSettings(app, fakeSettingsApi(false)); + + expect(app.ui.settings.definitions).toHaveLength(4); + expect(app.ui.settings.definitions[0]).toMatchObject({ + id: SIMPLE_SYRUP_SETTING_ID, + name: SIMPLE_SYRUP_SETTING_LABEL, + type: "boolean", + defaultValue: false, + tooltip: SIMPLE_SYRUP_SETTING_DESCRIPTION + }); + expect(SIMPLE_SYRUP_SETTING_DESCRIPTION).toContain("WD14 tagger"); + expect(SIMPLE_SYRUP_SETTING_DESCRIPTION).toContain("Ultralytics"); + expect(app.ui.settings.settings[0]?.value).toBe(false); + expect(app.ui.settings.definitions[1]).toMatchObject({ + id: QUANT_CACHE_SETTING_ID, + sortOrder: 321 + }); + expect(app.ui.settings.definitions[2]).toMatchObject({ + id: EXTERNAL_LLM_ENDPOINT_SETTING_ID, + sortOrder: 320 + }); + expect(app.ui.settings.definitions[3]).toMatchObject({ + id: EXTERNAL_LLM_API_KEY_SETTING_ID, + sortOrder: 319 + }); + }); + + it("saves downloadable-model setting changes", async () => { + const app = createFakeComfyApp(); + const refreshComboInNodes = vi.fn().mockResolvedValue(undefined); + app.refreshComboInNodes = refreshComboInNodes; + const saveSettings = vi + .fn() + .mockResolvedValue({ + show_downloadable_models: true, + quant_cache_limit_gib: 20 + }); + const api = fakeSettingsApi(false); + api.saveSettings = saveSettings; + + await registerSimpleSyrupSettings(app, api); + await app.ui.settings.definitions[0]?.onChange?.(true); + + expect(saveSettings).toHaveBeenCalledWith({ + show_downloadable_models: true, + quant_cache_limit_gib: 20 + }); + expect(app.ui.settings.settings[0]?.value).toBe(true); + expect(refreshComboInNodes).toHaveBeenCalledOnce(); + }); + + it("falls back to the default and warns when backend load fails", async () => { + const app = createFakeComfyApp(); + const logger = { warn: vi.fn() }; + const api = fakeSettingsApi(true); + api.getSettings = vi.fn().mockRejectedValue(new Error("offline")); + api.getExternalLLMSettings = vi + .fn() + .mockResolvedValue(defaultExternalLLMSettings()); + + await registerSimpleSyrupSettings(app, api, logger); + + expect(app.ui.settings.settings[0]?.value).toBe(true); + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("Could not load SimpleSyrup settings"), + expect.any(Error) + ); + }); + + it("restores the previous downloadable-model value when save fails", async () => { + const app = createFakeComfyApp(); + const logger = { warn: vi.fn() }; + const api = fakeSettingsApi(false); + api.saveSettings = vi + .fn() + .mockResolvedValueOnce({ + show_downloadable_models: true, + quant_cache_limit_gib: 20 + }) + .mockRejectedValueOnce(new Error("rejected")); + + await registerSimpleSyrupSettings(app, api, logger); + await app.ui.settings.definitions[0]?.onChange?.(true); + await app.ui.settings.definitions[0]?.onChange?.(false); + + expect(app.ui.settings.settings[0]?.value).toBe(true); + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("Could not save SimpleSyrup settings"), + expect.any(Error) + ); + }); + + it("keeps a saved setting when live model-choice refresh fails", async () => { + const app = createFakeComfyApp(); + const logger = { warn: vi.fn() }; + app.refreshComboInNodes = vi.fn().mockRejectedValue(new Error("offline")); + + await registerSimpleSyrupSettings(app, fakeSettingsApi(false), logger); + await app.ui.settings.definitions[0]?.onChange?.(true); + + expect(app.ui.settings.settings[0]?.value).toBe(true); + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("Could not refresh Comfy loader model choices"), + expect.any(Error) + ); + }); + + it("shows global quant cache usage and saves its GiB limit", async () => { + const app = createFakeComfyApp(); + const api = fakeSettingsApi(true); + const saveSettings = vi + .fn() + .mockResolvedValue({ + show_downloadable_models: true, + quant_cache_limit_gib: 30 + }); + api.saveSettings = saveSettings; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 1)); + const input = requiredInput(control, "input[type=number]"); + expect(input.value).toBe("20"); + expect(control.textContent).toContain("models/SyrupQuants"); + input.value = "30"; + requiredButton(control, "button").click(); + await flushPromises(); + + expect(saveSettings).toHaveBeenCalledWith({ + show_downloadable_models: true, + quant_cache_limit_gib: 30 + }); + expect(input.value).toBe("30"); + }); + + it("clears inactive quant artifacts and refreshes cache status", async () => { + const app = createFakeComfyApp(); + const api = fakeSettingsApi(true); + const clearQuantCache = vi.fn().mockResolvedValue({ + ...defaultQuantCacheStatus(), + removed_artifacts: 2, + removed_bytes: 1024 + }); + api.clearQuantCache = clearQuantCache; + + await registerSimpleSyrupSettings(app, api); + const control = renderSetting(getDefinition(app, 1)); + const clearButton = control.querySelectorAll("button")[1]; + if (!(clearButton instanceof HTMLButtonElement)) { + throw new Error("Expected quant cache clear button."); + } + clearButton.click(); + await flushPromises(); + + expect(clearQuantCache).toHaveBeenCalledOnce(); + expect(control.textContent).toContain("0.00 GiB used"); + }); + + it("restores the previous quant limit when backend saving fails", async () => { + const app = createFakeComfyApp(); + const logger = { warn: vi.fn() }; + const api = fakeSettingsApi(true); + api.saveSettings = vi.fn().mockRejectedValue(new Error("rejected")); + + await registerSimpleSyrupSettings(app, api, logger); + const control = renderSetting(getDefinition(app, 1)); + const input = requiredInput(control, "input[type=number]"); + input.value = "30"; + requiredButton(control, "button").click(); + await flushPromises(); + + expect(input.value).toBe("20"); + expect(control.textContent).toContain("was not saved"); + expect(logger.warn).toHaveBeenCalledWith( + expect.stringContaining("quant cache limit"), + expect.any(Error) + ); + }); +}); diff --git a/web/tests/settings/settingsTestSupport.ts b/web/tests/settings/settingsTestSupport.ts new file mode 100644 index 0000000..455d838 --- /dev/null +++ b/web/tests/settings/settingsTestSupport.ts @@ -0,0 +1,104 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { expect, vi } from "vitest"; + +import type { SimpleSyrupSettingsApi } from "../../src/settingsRegistration"; +import type { createFakeComfyApp } from "../support/testUtils"; + +export function fakeSettingsApi( + showDownloadableModels: boolean +): SimpleSyrupSettingsApi { + return { + getSettings: vi.fn().mockResolvedValue({ + show_downloadable_models: showDownloadableModels, + quant_cache_limit_gib: 20 + }), + saveSettings: vi.fn().mockImplementation((settings) => Promise.resolve(settings)), + getQuantCacheStatus: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + enforceQuantCacheLimit: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + clearQuantCache: vi.fn().mockResolvedValue(defaultQuantCacheStatus()), + getExternalLLMSettings: vi.fn().mockResolvedValue(defaultExternalLLMSettings()), + saveExternalLLMSettings: vi + .fn() + .mockImplementation((settings) => + Promise.resolve({ ...defaultExternalLLMSettings(), ...settings }) + ), + saveExternalLLMApiKey: vi.fn().mockResolvedValue(defaultExternalLLMSettings()) + }; +} + +export function defaultQuantCacheStatus() { + return { + path: "models/SyrupQuants", + usage_bytes: 0, + limit_bytes: 20 * 1024 ** 3, + artifact_count: 0, + active_artifact_count: 0 + }; +} + +export function defaultExternalLLMSettings() { + return { + base_url: "", + cached_models: [], + default_model: "", + has_api_key: false + }; +} + +export function configuredExternalLLMSettings() { + return { + base_url: "https://provider.example/v1", + cached_models: [], + default_model: "", + has_api_key: false + }; +} + +export function renderSetting(definition: { + type: "boolean" | "text" | (() => HTMLElement); +}): HTMLElement { + expect(typeof definition.type).toBe("function"); + return (definition.type as () => HTMLElement)(); +} + +export function getDefinition( + app: ReturnType, + index: number +) { + const definition = app.ui.settings.definitions[index]; + if (!definition) { + throw new Error(`Expected setting definition at index ${String(index)}.`); + } + return definition; +} + +export function requiredInput( + parent: ParentNode, + selector: string +): HTMLInputElement { + const input = parent.querySelector(selector); + if (!(input instanceof HTMLInputElement)) { + throw new Error(`Expected input for selector ${selector}.`); + } + return input; +} + +export function requiredButton( + parent: ParentNode, + selector: string +): HTMLButtonElement { + const button = parent.querySelector(selector); + if (!(button instanceof HTMLButtonElement)) { + throw new Error(`Expected button for selector ${selector}.`); + } + return button; +} + +export async function flushPromises(): Promise { + for (let index = 0; index < 6; index += 1) { + await Promise.resolve(); + } +} diff --git a/web/tests/testUtils.ts b/web/tests/support/testUtils.ts similarity index 99% rename from web/tests/testUtils.ts rename to web/tests/support/testUtils.ts index d4c7e76..43e70b6 100644 --- a/web/tests/testUtils.ts +++ b/web/tests/support/testUtils.ts @@ -8,7 +8,7 @@ import type { ComfySetting, ComfySettingDefinition, SettingValue -} from "../src/types"; +} from "../../src/types"; import { vi } from "vitest"; export interface FakeComfySettingsApi {