72 lines
2.3 KiB
Python
72 lines
2.3 KiB
Python
# 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 SimpleSyrup conditioning batch domain behavior."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import pytest
|
|
|
|
from simple_syrup.domain.conditioning_batch import (
|
|
ConditioningBatch,
|
|
batch_conditioning,
|
|
select_conditioning,
|
|
)
|
|
|
|
|
|
def test_conditioning_batch_requires_entries() -> None:
|
|
"""A batch must contain at least one selectable entry."""
|
|
|
|
with pytest.raises(ValueError, match="conditioning batch must contain"):
|
|
ConditioningBatch(())
|
|
|
|
|
|
def test_conditioning_batch_selects_by_index_with_last_entry_fallback() -> None:
|
|
"""Indexes beyond the batch length reuse the final entry."""
|
|
|
|
assert ConditioningBatch(("a",)).select(0) == "a"
|
|
assert ConditioningBatch(("a",)).select(5) == "a"
|
|
assert ConditioningBatch(("a", "b")).select(1) == "b"
|
|
assert ConditioningBatch(("a", "b")).select(5) == "b"
|
|
|
|
|
|
def test_conditioning_batch_rejects_negative_indexes() -> None:
|
|
"""Negative indexes are invalid for per-SEG selection."""
|
|
|
|
with pytest.raises(ValueError, match="conditioning batch index"):
|
|
ConditioningBatch(("a",)).select(-1)
|
|
|
|
|
|
def test_batch_conditioning_flattens_batches_and_normal_conditioning() -> None:
|
|
"""Mixed conditioning inputs become one ordered per-region batch."""
|
|
|
|
first = ConditioningBatch(("auto 1", "auto 2"))
|
|
hand = "hand 1"
|
|
second = ConditioningBatch(("auto 3",))
|
|
|
|
batch = batch_conditioning((first, hand, second))
|
|
|
|
assert batch.entries == ("auto 1", "auto 2", "hand 1", "auto 3")
|
|
|
|
|
|
def test_batch_conditioning_rejects_no_inputs() -> None:
|
|
"""At least one input is needed to build a conditioning batch."""
|
|
|
|
with pytest.raises(ValueError, match="one or more inputs"):
|
|
batch_conditioning(())
|
|
|
|
|
|
def test_select_conditioning_broadcasts_normal_conditioning() -> None:
|
|
"""Normal conditionings pass through unchanged for any valid index."""
|
|
|
|
conditioning = object()
|
|
|
|
assert select_conditioning(conditioning, 3) is conditioning
|
|
|
|
|
|
def test_select_conditioning_uses_batch_fallback() -> None:
|
|
"""Batch selection uses the same last-entry fallback policy."""
|
|
|
|
assert select_conditioning(ConditioningBatch(("a", "b")), 5) == "b"
|