kneifftools / tests /test_utils_helpers.py
kneiff's picture
chore(history)!: publish anonymous repository root
2857cf3
Raw
History Blame Contribute Delete
18.9 kB
from __future__ import annotations
import logging
from pathlib import Path
import subprocess
import sys
from PIL import Image, ImageChops
import pytest
from rich.console import Console
import kneiff.utils as ut
import kneiff.utils.config as config_utils
import kneiff.utils.env as env_utils
import kneiff.utils.image.caption.io as caption_io
import kneiff.utils.image.caption.tags as caption_tags
import kneiff.utils.image.caption.text as caption_text
import kneiff.utils.image.discovery as image_discovery
import kneiff.utils.image.encoding as image_encoding
import kneiff.utils.image.grid as image_grid
import kneiff.utils.image.sidecars as image_sidecars
import kneiff.utils.logging as logging_utils
import kneiff.utils.tablefmt as tablefmt_utils
import kneiff.utils.text as text_utils
import kneiff.utils.yaml as yaml_utils
def test_utils_facade_reexports_shared_helpers() -> None:
"""Keep the public utility facade aligned with the shared helper modules."""
assert ut.read_dotenv is env_utils.read_dotenv
assert ut.read_json_mapping is config_utils.read_json_mapping
assert ut.read_yaml_document is config_utils.read_yaml_document
assert ut.dump_yaml is yaml_utils.dump_yaml
assert ut.export_yaml is yaml_utils.export_yaml
assert ut.load_yaml is yaml_utils.load_yaml
assert ut.YamlError is yaml_utils.YamlError
assert ut.optional_text is config_utils.optional_text
assert ut.text_tuple is config_utils.text_tuple
assert ut.format_text_table is tablefmt_utils.format_text_table
assert ut.split_comma_separated_values is text_utils.split_comma_separated_values
def test_utils_facade_keeps_rich_lazy_for_non_cli_imports() -> None:
"""Protect Space imports that only need config helpers from CLI-only Rich."""
script = (
"import sys\n"
"import kneiff.utils\n"
"raise SystemExit(1 if 'rich' in sys.modules else 0)\n"
)
result = subprocess.run(
[sys.executable, "-c", script],
capture_output=True,
text=True,
check=False,
)
assert result.returncode == 0, result.stdout + result.stderr
def test_stdlib_deep_mapping_helpers() -> None:
"""Cover nested mapping helpers exposed through the utility facade."""
data: dict[object, object] = {"a": {"b": 1}}
assert ut.deep_get(data, ("a", "b")) == 1
assert ut.deep_get(data, ("a", "missing"), default="fallback") == "fallback"
assert ut.deep_get(data, (), default="fallback") == "fallback"
ut.deep_set(data, ("a", "c"), 2)
assert data == {"a": {"b": 1, "c": 2}}
assert ut.deep_right_merge(
{"a": {"b": 1}, "keep": True},
{"a": {"c": 2}},
) == {"a": {"b": 1, "c": 2}, "keep": True}
def test_read_dotenv_parses_exports_quotes_and_comments(tmp_path: Path) -> None:
env_path = tmp_path / ".env.shared"
env_path.write_text(
"""
# ignored
export KNF_WORKERS=' 4 '
EMPTY=
MODEL_NAME = "local model"
invalid line
""",
encoding="utf-8",
)
values = env_utils.read_dotenv(env_path)
assert values == {
"KNF_WORKERS": "4",
"EMPTY": "",
"MODEL_NAME": "local model",
}
def test_read_dotenv_value_returns_non_empty_value(tmp_path: Path) -> None:
env_path = tmp_path / ".env.shared"
env_path.write_text("KNF_WORKERS= 4 \nEMPTY=\n", encoding="utf-8")
assert env_utils.read_dotenv_value(env_path, "KNF_WORKERS") == "4"
assert env_utils.read_dotenv_value(env_path, "EMPTY") is None
assert env_utils.read_dotenv_value(env_path, "MISSING") is None
assert env_utils.read_dotenv_value(tmp_path / "missing.env", "KNF_WORKERS") is None
with pytest.raises(RuntimeError, match=r"\.env file not found"):
env_utils.read_dotenv_value(
tmp_path / "missing.env",
"KNF_WORKERS",
require_file=True,
)
def test_first_env_value_prefers_process_env(
monkeypatch: pytest.MonkeyPatch,
) -> None:
file_values = {"KNF_WORKERS": "2", "ALT_WORKERS": "3"}
monkeypatch.delenv("KNF_WORKERS", raising=False)
monkeypatch.delenv("ALT_WORKERS", raising=False)
assert (
env_utils.first_env_value(
["KNF_WORKERS", "ALT_WORKERS"], dotenv_values=file_values
)
== "2"
)
monkeypatch.setenv("KNF_WORKERS", "5")
assert (
env_utils.first_env_value(
["KNF_WORKERS", "ALT_WORKERS"], dotenv_values=file_values
)
== "5"
)
def test_parse_positive_int_defaults_and_rejects_invalid_values() -> None:
assert env_utils.parse_positive_int(None, name="KNF_WORKERS", source="env") == 1
assert env_utils.parse_positive_int(" 3 ", name="KNF_WORKERS", source="env") == 3
with pytest.raises(ValueError, match="KNF_WORKERS"):
env_utils.parse_positive_int("0", name="KNF_WORKERS", source="env")
def test_config_helpers_load_yaml_and_validate_scalar_nodes(tmp_path: Path) -> None:
config_path = tmp_path / "config.yaml"
config_path.write_text("enabled: true\ncount: 3\n", encoding="utf-8")
document = config_utils.read_yaml_document(
config_path,
dependency_error="py-yaml12 is required.",
)
mapping = config_utils.require_mapping(
document, message="config must be a mapping."
)
assert config_utils.require_bool(mapping["enabled"], name="enabled") is True
assert config_utils.require_positive_int(mapping["count"], name="count") == 3
assert config_utils.require_exact_text("subset", name="subset") == "subset"
assert config_utils.optional_text(" caption ", name="caption") == "caption"
assert config_utils.text_tuple(
[" tag ", "", "tag", "other"],
name="tags",
) == ("tag", "other")
assert config_utils.require_int(5, name="seed") == 5
assert (
config_utils.optional_mapping(None, message="optional must be a mapping.")
is None
)
with pytest.raises(ValueError, match="enabled must be true or false"):
config_utils.require_bool("yes", name="enabled")
with pytest.raises(ValueError, match="count must be a positive integer"):
config_utils.require_positive_int(True, name="count")
with pytest.raises(ValueError, match="seed must be an integer"):
config_utils.require_int("5", name="seed")
with pytest.raises(ValueError, match="subset must be a non-empty exact string"):
config_utils.require_exact_text(" subset ", name="subset")
def test_yaml_helpers_export_readable_multiline_strings(tmp_path: Path) -> None:
output_path = tmp_path / "nested" / "manifest.yaml"
yaml_utils.export_yaml(
{
"caption": "line 1\nline 2\n",
"path": Path("SOURCE/scene.png"),
"tags": {"beta", "alpha"},
},
output_path,
)
text = output_path.read_text(encoding="utf-8")
assert "caption: |2" in text
assert "line 1" in text
assert "line 2" in text
assert yaml_utils.load_yaml(output_path) == {
"caption": "line 1\nline 2\n",
"path": "SOURCE/scene.png",
"tags": ["alpha", "beta"],
}
def test_yaml_helpers_load_empty_text_as_none() -> None:
assert yaml_utils.load_yaml_text("") is None
def test_read_yaml_document_rejects_duplicate_top_level_keys(tmp_path: Path) -> None:
config_path = tmp_path / "config.yaml"
config_path.write_text("enabled: true\nenabled: false\n", encoding="utf-8")
with pytest.raises(ValueError, match="Duplicate YAML key 'enabled'"):
config_utils.read_yaml_document(
config_path,
dependency_error="py-yaml12 is required.",
)
def test_read_yaml_document_rejects_duplicate_nested_keys(tmp_path: Path) -> None:
config_path = tmp_path / "config.yaml"
config_path.write_text(
"""
training:
simpletuner:
trainer:
caption_dropout_probability: 0.0
caption_dropout_probability: 0.05
""",
encoding="utf-8",
)
with pytest.raises(
ValueError,
match="Duplicate YAML key 'caption_dropout_probability'",
):
config_utils.read_yaml_document(
config_path,
dependency_error="py-yaml12 is required.",
)
def test_config_helpers_load_json_file_shapes(tmp_path: Path) -> None:
mapping_path = tmp_path / "mapping.json"
sequence_path = tmp_path / "sequence.json"
mapping_path.write_text('{"alpha": 1}', encoding="utf-8")
sequence_path.write_text('[{"id": "a"}, {"id": "b"}]', encoding="utf-8")
assert config_utils.read_json_mapping(mapping_path) == {"alpha": 1}
assert config_utils.read_json_sequence(sequence_path) == [
{"id": "a"},
{"id": "b"},
]
assert config_utils.read_json_mapping_sequence(sequence_path) == [
{"id": "a"},
{"id": "b"},
]
def test_config_helpers_reject_unexpected_json_shapes(tmp_path: Path) -> None:
mapping_path = tmp_path / "mapping.json"
sequence_path = tmp_path / "sequence.json"
invalid_path = tmp_path / "invalid.json"
mapping_path.write_text('{"alpha": 1}', encoding="utf-8")
sequence_path.write_text('["not-an-object"]', encoding="utf-8")
invalid_path.write_text("{", encoding="utf-8")
with pytest.raises(ValueError, match="Expected JSON array"):
config_utils.read_json_sequence(mapping_path)
with pytest.raises(ValueError, match="Expected JSON object entries"):
config_utils.read_json_mapping_sequence(sequence_path)
with pytest.raises(ValueError, match="Broken JSON"):
config_utils.read_json_document(
invalid_path,
decode_error_message="Broken JSON",
)
def test_tablefmt_styles_plain_text_table_for_console() -> None:
rendered = tablefmt_utils.format_text_table(
headers=("Subset", "Images", "Train %"),
rows=[("fullbody", 30, "100.0%")],
right_align_headers={"Images", "Train %"},
)
styled = tablefmt_utils.style_text_table_for_console(
rendered,
column_styles={"Subset": "bold", "Train %": "green"},
)
assert styled.plain == rendered
assert styled.spans
plain_console = Console(
force_terminal=True,
color_system="truecolor",
width=120,
highlight=False,
)
with plain_console.capture() as capture:
plain_console.print(rendered)
plain_output = capture.get()
styled_console = Console(
force_terminal=True,
color_system="truecolor",
width=120,
highlight=False,
)
with styled_console.capture() as capture:
styled_console.print(styled)
styled_output = capture.get()
assert "\x1b[" not in plain_output
assert "\x1b[" in styled_output
assert "fullbody" in styled_output
def test_tablefmt_styles_exact_cell_values_for_console() -> None:
rendered = tablefmt_utils.format_text_table(
headers=("Status", "Images"),
rows=[("KEPT", 2), ("KICKED", 1)],
right_align_headers={"Images"},
)
styled = tablefmt_utils.style_text_table_for_console(
rendered,
column_styles={"Status": "green"},
cell_styles={"KICKED": "bold bright_red"},
)
assert styled.plain == rendered
kicked_start = rendered.index("KICKED")
kicked_spans = [
span
for span in styled.spans
if span.start <= kicked_start and span.end >= kicked_start + len("KICKED")
]
assert any("bright_red" in str(span.style) for span in kicked_spans)
def test_image_grid_writes_grouped_contact_sheet_with_labels(tmp_path: Path) -> None:
red_path = tmp_path / "red.png"
blue_path = tmp_path / "blue.png"
output_path = tmp_path / "grid.png"
Image.new("RGB", (12, 8), color="red").save(red_path)
Image.new("RGB", (8, 12), color="blue").save(blue_path)
written_path = image_grid.write_grouped_contact_sheet(
[
image_grid.ContactSheetGroup(
label="fullbody",
images=(
image_grid.GridImage(red_path),
image_grid.GridImage(blue_path),
),
)
],
output_path,
thumbnail_size=16,
columns=2,
)
assert written_path == output_path
with Image.open(output_path) as rendered:
assert rendered.size[0] >= 16 * 2
assert rendered.size[1] > 16
header = rendered.crop((0, 0, rendered.width, 16))
blank_header = Image.new("RGB", header.size, color="white")
assert ImageChops.difference(header, blank_header).getbbox() is not None
def test_image_grid_wrap_text_block_caps_lines_with_ellipsis() -> None:
font = image_grid._load_font(24)
wrapped = image_grid._wrap_text_block(
(
"alpha\n"
"This validation prompt text is long enough to wrap across many lines "
"and should end with an ellipsis once the limit is reached."
),
font=font,
max_width=120,
max_lines=5,
)
assert len(wrapped) == 5
assert wrapped[0] == "alpha"
assert wrapped[-1].endswith("...")
def test_image_grid_prompt_label_limit_allows_ten_wrapped_lines() -> None:
font = image_grid._load_font(24)
wrapped = image_grid._wrap_text_block(
" ".join("word" for _ in range(100)),
font=font,
max_width=100,
max_lines=image_grid.DEFAULT_PROMPT_LABEL_MAX_LINES,
)
assert image_grid.DEFAULT_PROMPT_LABEL_MAX_LINES == 10
assert len(wrapped) == 10
assert wrapped[-1].endswith("...")
def test_image_grid_writes_matrix_contact_sheet_with_footer_labels(
tmp_path: Path,
) -> None:
image_path = tmp_path / "white.png"
output_path = tmp_path / "matrix-grid.png"
Image.new("RGB", (16, 16), color="white").save(image_path)
written_path = image_grid.write_matrix_contact_sheet(
row_labels=("step 1",),
column_labels=("alpha_prompt",),
cells={("step 1", "alpha_prompt"): image_grid.GridImage(image_path)},
output_path=output_path,
thumbnail_size=32,
column_display_text={
"alpha_prompt": (
"alpha_prompt\n"
"A long validation prompt that wraps across several lines so the "
"footer label region is visible in the rendered sheet."
)
},
repeat_column_labels_at_bottom=True,
column_label_max_lines=5,
)
assert written_path == output_path
with Image.open(output_path) as rendered:
top = rendered.crop((120, 0, rendered.width, min(100, rendered.height)))
bottom = rendered.crop(
(120, max(0, rendered.height - 100), rendered.width, rendered.height)
)
blank_top = Image.new("RGB", top.size, color="white")
blank_bottom = Image.new("RGB", bottom.size, color="white")
assert ImageChops.difference(top, blank_top).getbbox() is not None
assert ImageChops.difference(bottom, blank_bottom).getbbox() is not None
def test_split_comma_separated_values_ignores_empty_tokens() -> None:
assert text_utils.split_comma_separated_values(" alpha, beta ,, gamma ") == [
"alpha",
"beta",
"gamma",
]
def test_iter_image_files_returns_sorted_supported_images(tmp_path: Path) -> None:
(tmp_path / "b.JPG").write_bytes(b"image")
(tmp_path / "a.png").write_bytes(b"image")
(tmp_path / "notes.txt").write_text("skip", encoding="utf-8")
nested = tmp_path / "nested"
nested.mkdir()
(nested / "c.webp").write_bytes(b"image")
assert image_discovery.iter_image_files(tmp_path, recursive=False) == [
tmp_path / "a.png",
tmp_path / "b.JPG",
]
assert image_discovery.iter_image_files(tmp_path) == [
tmp_path / "a.png",
tmp_path / "b.JPG",
nested / "c.webp",
]
def test_comma_sidecars_read_and_write_without_overwrite(tmp_path: Path) -> None:
image_path = tmp_path / "scene.png"
image_path.write_bytes(b"image")
sidecar = image_path.with_suffix(".txt")
assert image_sidecars.read_comma_sidecar(image_path, ".txt") == []
sidecar.write_text(" alpha, beta , , gamma\n", encoding="utf-8")
assert image_sidecars.read_comma_sidecar(image_path, ".txt") == [
"alpha",
"beta",
"gamma",
]
image_sidecars.write_text_sidecar(sidecar, "replacement", overwrite=False)
assert sidecar.read_text(encoding="utf-8") == " alpha, beta , , gamma\n"
image_sidecars.write_text_sidecar(sidecar, "replacement", overwrite=True)
assert sidecar.read_text(encoding="utf-8") == "replacement\n"
def test_caption_sidecar_io_builds_output_paths() -> None:
io = caption_io.CaptionSidecarIO(out_suffix=".cap.txt")
assert io.output_path(Path("frame.png")) == Path("frame.cap.txt")
def test_caption_tag_filtering_can_match_blip_aliases() -> None:
tags = [" alpha_tag ", "Beta", "alpha tag", "gamma"]
assert caption_tags.filter_caption_tags(
tags,
ignore_tags=["beta"],
replace_underscores=True,
dedupe=True,
) == ["alpha_tag", "gamma"]
assert caption_tags.filter_caption_tags(
tags,
ignore_tags=["beta"],
) == ["alpha_tag", "alpha tag", "gamma"]
def test_caption_text_rules_preserve_backend_specific_ban_behavior() -> None:
raw_caption = '"masterpiece Rook_Kaefer standing in best quality light"'
blip_rules = caption_text.build_caption_text_rules(
must_include=["felkin"],
ban_phrases=None,
max_words=6,
strip_banned=True,
)
server_rules = caption_text.build_caption_text_rules(
must_include=["felkin"],
ban_phrases=None,
max_words=6,
strip_banned=False,
)
assert blip_rules.apply(raw_caption) == "felkin Rook_Kaefer standing in light"
assert (
server_rules.apply(raw_caption)
== "felkin masterpiece Rook_Kaefer standing in best"
)
def test_image_to_jpeg_data_url_downscales_and_encodes(tmp_path: Path) -> None:
image_path = tmp_path / "source.png"
Image.new("RGB", (12, 6), color=(20, 30, 40)).save(image_path)
assert (
image_encoding.image_to_jpeg_data_url(image_path, max_side=4).split(",", 1)[0]
== "data:image/jpeg;base64"
)
def test_get_logger_builds_named_logger() -> None:
assert logging_utils.get_logger("kneiff.tests").name == "kneiff.tests"
def test_configure_logging_sets_root_level() -> None:
root_logger = logging.getLogger()
original_level = root_logger.level
try:
logging_utils.configure_logging(level="DEBUG")
assert root_logger.level == logging.DEBUG
finally:
root_logger.setLevel(original_level)
def test_configure_logging_rejects_unknown_level() -> None:
with pytest.raises(ValueError, match="Unknown logging level"):
logging_utils.configure_logging(level="verbose-ish")