File size: 3,769 Bytes
cb617a4
 
 
 
 
 
 
 
 
 
 
 
 
 
5fff6c4
cb617a4
5fff6c4
cb617a4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5fff6c4
 
 
 
 
 
 
 
cb617a4
 
 
 
 
 
 
 
 
 
 
 
 
5fff6c4
 
 
cb617a4
 
 
 
 
 
 
5fff6c4
cb617a4
 
 
 
 
 
 
 
 
 
5fff6c4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
cb617a4
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
"""Unit tests for Krea profile and image metadata helpers."""

from __future__ import annotations

import tempfile
import unittest
from pathlib import Path

from PIL import Image

from settings_utils import (
    APP_ID,
    build_settings,
    extract_image_settings,
    normalize_custom_loras,
    parse_settings_text,
    validate_custom_lora,
    write_png_metadata,
)


class SettingsTests(unittest.TestCase):
    def settings(self):
        return build_settings(
            mode="edit",
            prompt="",
            edit_prompt="change the coat to blue",
            width=1024,
            height=768,
            target_megapixels=1.4,
            grounding_px=768,
            ref_boost=1.0,
            ref_boost_a=1.0,
            steps=8,
            cfg=1.0,
            sampler_name="euler",
            scheduler="beta",
            seed=2,
            randomize_seed=False,
            gen_budget=0,
            effective_seed=2,
            base_model="pornmasterKrea2_v2TurboInt8.safetensors",
            catalog_loras=[{"hf_filename": "slider.safetensors", "weight": 0.8}],
            custom_loras=[{
                "repo_id": "org/repo",
                "filename": "custom.safetensors",
                "revision": "main",
                "weight": 0.5,
            }],
        )

    def test_png_metadata_round_trip(self):
        settings = self.settings()
        with tempfile.TemporaryDirectory() as directory:
            source = Path(directory) / "source.png"
            destination = Path(directory) / "output.png"
            Image.new("RGB", (8, 8), (1, 2, 3)).save(source)
            write_png_metadata(source, destination, settings)
            loaded, warnings = extract_image_settings(destination)
        self.assertEqual(loaded["app"], APP_ID)
        self.assertEqual(loaded["edit_prompt"], "change the coat to blue")
        self.assertEqual(loaded["ref_boost"], 1.0)
        self.assertEqual(loaded["base_model"], "pornmasterKrea2_v2TurboInt8.safetensors")
        self.assertEqual(loaded["catalog_loras"][0]["weight"], 0.8)
        self.assertEqual(loaded["custom_loras"][0]["repo_id"], "org/repo")
        self.assertEqual(warnings, [])

    def test_json_profile_round_trip(self):
        parsed, warnings = parse_settings_text(__import__("json").dumps(self.settings()))
        self.assertEqual(warnings, [])
        self.assertEqual(parsed["mode"], "edit")
        self.assertEqual(parsed["grounding_px"], 768)
        self.assertEqual(parsed["custom_loras"][0]["filename"], "custom.safetensors")

    def test_a1111_parameter_text(self):
        parsed, warnings = parse_settings_text(
            "a portrait\nSteps: 8, CFG scale: 1, Sampler: euler, Seed: 7, Size: 512x768"
        )
        self.assertEqual(warnings, [])
        self.assertEqual(parsed["prompt"], "a portrait")
        self.assertEqual(parsed["width"], 512)
        self.assertEqual(parsed["effective_seed"], 7)

    def test_custom_lora_validation_blocks_unsafe_rows(self):
        valid = validate_custom_lora(["org/repo", "sub/style.safetensors", "main", 0.8])
        self.assertEqual(valid["filename"], "sub/style.safetensors")
        normalized, warnings = normalize_custom_loras([
            ["", "", "", 0],
            ["org/repo", "../escape.safetensors", "", 1],
        ])
        self.assertEqual(normalized, [])
        self.assertEqual(len(warnings), 1)
        with self.assertRaises(ValueError):
            validate_custom_lora({
                "repo_id": "not-a-repo",
                "filename": "style.safetensors",
                "weight": 0.5,
            })


if __name__ == "__main__":
    unittest.main()