import ast import base64 import json import struct import threading import time import unittest import zlib from io import BytesIO from pathlib import Path from unittest.mock import patch from PIL import Image, features from provenance import ENABLED_ADAPTER_SPECS from i2i_contract import ( BASE_EDIT_MODE_ID, ContractError, MAX_IMAGE_BYTES, MAX_INPUT_IMAGES, MAX_PROMPT_CHARS, MAX_REQUEST_JSON_CHARS, MAX_SEED, MAX_TOTAL_IMAGE_BYTES, MAX_TOTAL_IMAGE_PIXELS, MIN_INPUT_IMAGES, SingleResidentAdapterManager, decode_validated_images, encode_png_result, execute_validated_edit, output_dimensions, parse_example_index, serialize_png_candidate, validate_edit_request, ) ALLOWED_EDIT_MODES = {BASE_EDIT_MODE_ID, "Anime-V2"} ROOT = Path(__file__).resolve().parents[1] def image_bytes(fmt="PNG", size=(128, 96), color=(20, 40, 60)): buffer = BytesIO() Image.new("RGB", size, color).save(buffer, format=fmt) return buffer.getvalue() def data_url(data, mime="image/png"): return f"data:{mime};base64,{base64.b64encode(data).decode('ascii')}" def request_args(**overrides): values = { "images_b64_json": json.dumps([data_url(image_bytes())]), "prompt": " improve the lighting ", "lora_adapter": BASE_EDIT_MODE_ID, "seed": 0, "randomize_seed": True, "guidance_scale": 1.0, "steps": 4, "allowed_edit_modes": ALLOWED_EDIT_MODES, } values.update(overrides) return values def png_header_only(width, height): signature = b"\x89PNG\r\n\x1a\n" payload = struct.pack(">IIBBBBB", width, height, 8, 2, 0, 0, 0) ihdr = b"IHDR" + payload return signature + struct.pack(">I", len(payload)) + ihdr + struct.pack(">I", zlib.crc32(ihdr)) class ContractTests(unittest.TestCase): def assert_code(self, expected, **overrides): with self.assertRaises(ContractError) as caught: validate_edit_request(**request_args(**overrides)) self.assertEqual(caught.exception.code, expected) def test_valid_png_request_is_normalized_and_keeps_bytes(self): request = validate_edit_request(**request_args()) self.assertEqual(request.prompt, "improve the lighting") self.assertEqual(request.edit_mode, BASE_EDIT_MODE_ID) self.assertEqual(request.images[0].mime, "image/png") self.assertIsInstance(request.images[0].data, bytes) self.assertEqual((request.images[0].width, request.images[0].height), (128, 96)) decoded = decode_validated_images(request.images) self.assertEqual(decoded[0].size, (128, 96)) self.assertEqual(output_dimensions(request.images[0]), (1024, 768)) def test_one_or_two_images_are_accepted(self): first = data_url(image_bytes(color=(1, 2, 3))) second = data_url(image_bytes(color=(4, 5, 6))) one = validate_edit_request(**request_args(images_b64_json=json.dumps([first]))) two = validate_edit_request(**request_args(images_b64_json=json.dumps([first, second]))) self.assertEqual(len(one.images), MIN_INPUT_IMAGES) self.assertEqual(len(two.images), MAX_INPUT_IMAGES) self.assertEqual([image.width for image in two.images], [128, 128]) def test_jpeg_and_webp_are_accepted(self): jpeg = json.dumps([data_url(image_bytes("JPEG"), "image/jpeg")]) self.assertEqual(validate_edit_request(**request_args(images_b64_json=jpeg)).images[0].mime, "image/jpeg") if features.check("webp"): webp = json.dumps([data_url(image_bytes("WEBP"), "image/webp")]) self.assertEqual(validate_edit_request(**request_args(images_b64_json=webp)).images[0].mime, "image/webp") def test_json_and_image_count_are_strict(self): self.assert_code("I2I_BAD_REQUEST", images_b64_json="not-json") self.assert_code("I2I_BAD_REQUEST", images_b64_json="{}") self.assert_code("I2I_BAD_REQUEST", images_b64_json="[]") three = json.dumps([data_url(image_bytes())] * 3) self.assert_code("I2I_BAD_REQUEST", images_b64_json=three) self.assert_code("I2I_BAD_REQUEST", images_b64_json=json.dumps([123])) def test_mime_base64_and_magic_are_strict(self): self.assert_code("I2I_UNSUPPORTED_TYPE", images_b64_json=json.dumps(["data:image/gif;base64,AAAA"])) self.assert_code("I2I_IMAGE_DECODE_FAILED", images_b64_json=json.dumps(["data:image/png;base64,***"])) mismatch = json.dumps([data_url(image_bytes("JPEG"), "image/png")]) self.assert_code("I2I_UNSUPPORTED_TYPE", images_b64_json=mismatch) corrupt = json.dumps([data_url(b"\x89PNG\r\n\x1a\ncorrupt")]) self.assert_code("I2I_IMAGE_DECODE_FAILED", images_b64_json=corrupt) def test_encoded_size_is_rejected_before_decode(self): too_long = "A" * (4 * ((MAX_IMAGE_BYTES + 2) // 3) + 1) self.assert_code( "I2I_FILE_TOO_LARGE", images_b64_json=json.dumps([f"data:image/png;base64,{too_long}"]), ) self.assert_code("I2I_FILE_TOO_LARGE", images_b64_json=" " * (MAX_REQUEST_JSON_CHARS + 1)) def test_total_byte_and_pixel_limits_are_enforced(self): raw = image_bytes() two = json.dumps([data_url(raw), data_url(raw)]) with patch("i2i_contract.MAX_TOTAL_IMAGE_BYTES", len(raw) * 2 - 1): self.assert_code("I2I_FILE_TOO_LARGE", images_b64_json=two) with patch("i2i_contract.MAX_TOTAL_IMAGE_PIXELS", (128 * 96 * 2) - 1): self.assert_code("I2I_IMAGE_DIMENSIONS_INVALID", images_b64_json=two) self.assertEqual(MAX_TOTAL_IMAGE_BYTES, MAX_IMAGE_BYTES * MAX_INPUT_IMAGES) self.assertEqual(MAX_TOTAL_IMAGE_PIXELS, 16_000_000 * MAX_INPUT_IMAGES) def test_dimensions_are_bounded(self): small = json.dumps([data_url(image_bytes(size=(63, 128)))]) self.assert_code("I2I_IMAGE_DIMENSIONS_INVALID", images_b64_json=small) oversized = json.dumps([data_url(png_header_only(4001, 4000))]) self.assert_code("I2I_IMAGE_DIMENSIONS_INVALID", images_b64_json=oversized) extreme = json.dumps([data_url(image_bytes(size=(64, 9000)))]) self.assert_code("I2I_IMAGE_DIMENSIONS_INVALID", images_b64_json=extreme) def test_prompt_edit_mode_and_numbers_are_bounded(self): self.assert_code("I2I_PROMPT_REQUIRED", prompt=" ") self.assert_code("I2I_PROMPT_TOO_LONG", prompt="x" * (MAX_PROMPT_CHARS + 1)) self.assert_code("I2I_LORA_NOT_ALLOWED", lora_adapter="unknown") self.assertEqual( validate_edit_request(**request_args(lora_adapter="Anime-V2")).edit_mode, "Anime-V2", ) self.assert_code("I2I_BAD_REQUEST", seed=True) self.assert_code("I2I_BAD_REQUEST", seed=-1) self.assert_code("I2I_BAD_REQUEST", seed=MAX_SEED + 1) self.assert_code("I2I_BAD_REQUEST", randomize_seed=1) self.assert_code("I2I_BAD_REQUEST", guidance_scale=float("nan")) self.assert_code("I2I_BAD_REQUEST", guidance_scale=float("inf")) self.assert_code("I2I_BAD_REQUEST", guidance_scale=0.99) self.assert_code("I2I_BAD_REQUEST", guidance_scale=10.01) self.assert_code("I2I_BAD_REQUEST", steps=True) self.assert_code("I2I_BAD_REQUEST", steps=0) self.assert_code("I2I_BAD_REQUEST", steps=51) def test_invalid_input_never_calls_runner(self): calls = 0 def runner(_request): nonlocal calls calls += 1 return image_bytes(), 1 bad_cases = [ {"images_b64_json": "[]"}, {"prompt": ""}, {"lora_adapter": "unknown"}, {"steps": 999}, ] for case in bad_cases: with self.assertRaises(ContractError): execute_validated_edit(**request_args(**case), runner=runner) self.assertEqual(calls, 0) def test_valid_runner_result_is_one_png_and_seed(self): calls = 0 def runner(_request): nonlocal calls calls += 1 return image_bytes(), 123 result = execute_validated_edit(**request_args(), runner=runner) self.assertEqual(calls, 1) self.assertEqual(result["seed"], 123) self.assertTrue(result["image"].startswith("data:image/png;base64,")) def test_pipeline_style_lock_serializes_mutation(self): lock = threading.Lock() barrier = threading.Barrier(3) active = 0 max_active = 0 def mutate(): nonlocal active, max_active barrier.wait() with lock: active += 1 max_active = max(max_active, active) time.sleep(0.01) active -= 1 workers = [threading.Thread(target=mutate) for _ in range(2)] for worker in workers: worker.start() barrier.wait() for worker in workers: worker.join() self.assertEqual(max_active, 1) def test_single_resident_adapter_state_machine(self): events = [] manager = SingleResidentAdapterManager( adapter_modes={"Anime-V2", "Style-Transfer"}, load_adapter=lambda mode: events.append(("load", mode)), activate_adapter=lambda mode: events.append(("activate", mode)), unload_adapters=lambda: events.append(("unload", None)), ) manager.activate_mode(BASE_EDIT_MODE_ID) self.assertEqual(events, []) manager.activate_mode("Anime-V2") manager.activate_mode("Anime-V2") self.assertEqual(manager.resident_adapter_mode, "Anime-V2") self.assertFalse(manager.dirty) manager.activate_mode("Style-Transfer") manager.activate_mode(BASE_EDIT_MODE_ID) self.assertEqual( events, [ ("load", "Anime-V2"), ("activate", "Anime-V2"), ("unload", None), ("load", "Style-Transfer"), ("activate", "Style-Transfer"), ("unload", None), ], ) self.assertIsNone(manager.resident_adapter_mode) self.assertFalse(manager.dirty) with self.assertRaises(ContractError) as caught: manager.activate_mode("unknown") self.assertEqual(caught.exception.code, "I2I_LORA_NOT_ALLOWED") def test_partial_adapter_load_is_cleaned_before_retry(self): events = [] def load(mode): events.append(("load", mode)) if mode == "Style-Transfer": raise RuntimeError("synthetic load failure") manager = SingleResidentAdapterManager( adapter_modes={"Anime-V2", "Style-Transfer"}, load_adapter=load, activate_adapter=lambda mode: events.append(("activate", mode)), unload_adapters=lambda: events.append(("unload", None)), ) manager.activate_mode("Anime-V2") with self.assertRaises(RuntimeError): manager.activate_mode("Style-Transfer") self.assertIsNone(manager.resident_adapter_mode) self.assertFalse(manager.dirty) self.assertEqual(events[-3:], [("unload", None), ("load", "Style-Transfer"), ("unload", None)]) manager.activate_mode(BASE_EDIT_MODE_ID) self.assertEqual(events[-1], ("unload", None)) def test_adapter_activation_failure_is_cleaned_before_next_mode(self): events = [] fail_activation_once = True def activate(mode): nonlocal fail_activation_once events.append(("activate", mode)) if fail_activation_once: fail_activation_once = False raise RuntimeError("synthetic activation failure") manager = SingleResidentAdapterManager( adapter_modes={"Anime-V2", "Style-Transfer"}, load_adapter=lambda mode: events.append(("load", mode)), activate_adapter=activate, unload_adapters=lambda: events.append(("unload", None)), ) with self.assertRaises(RuntimeError): manager.activate_mode("Anime-V2") self.assertIsNone(manager.resident_adapter_mode) self.assertFalse(manager.dirty) self.assertEqual( events, [ ("load", "Anime-V2"), ("activate", "Anime-V2"), ("unload", None), ], ) manager.activate_mode("Style-Transfer") self.assertEqual(manager.resident_adapter_mode, "Style-Transfer") self.assertFalse(manager.dirty) self.assertEqual( events[-2:], [("load", "Style-Transfer"), ("activate", "Style-Transfer")], ) def test_failed_cleanup_stays_dirty_and_is_retried_before_base_mode(self): events = [] fail_load_once = True fail_unload_once = True def load(mode): nonlocal fail_load_once events.append(("load", mode)) if fail_load_once: fail_load_once = False raise RuntimeError("synthetic load failure") def unload(): nonlocal fail_unload_once events.append(("unload", None)) if fail_unload_once: fail_unload_once = False raise RuntimeError("synthetic cleanup failure") manager = SingleResidentAdapterManager( adapter_modes={"Anime-V2"}, load_adapter=load, activate_adapter=lambda mode: events.append(("activate", mode)), unload_adapters=unload, ) with self.assertRaises(RuntimeError): manager.activate_mode("Anime-V2") self.assertIsNone(manager.resident_adapter_mode) self.assertTrue(manager.dirty) manager.activate_mode(BASE_EDIT_MODE_ID) self.assertIsNone(manager.resident_adapter_mode) self.assertFalse(manager.dirty) self.assertEqual( events, [ ("load", "Anime-V2"), ("unload", None), ("unload", None), ], ) def test_switch_unload_failure_is_retried_before_loading_new_mode(self): events = [] unload_calls = 0 def unload(): nonlocal unload_calls unload_calls += 1 events.append(("unload", None)) if unload_calls == 1: raise RuntimeError("synthetic unload failure") manager = SingleResidentAdapterManager( adapter_modes={"Anime-V2", "Style-Transfer"}, load_adapter=lambda mode: events.append(("load", mode)), activate_adapter=lambda mode: events.append(("activate", mode)), unload_adapters=unload, ) manager.activate_mode("Anime-V2") with self.assertRaises(RuntimeError): manager.activate_mode("Style-Transfer") self.assertIsNone(manager.resident_adapter_mode) self.assertTrue(manager.dirty) self.assertNotIn(("load", "Style-Transfer"), events) manager.activate_mode("Style-Transfer") self.assertEqual(manager.resident_adapter_mode, "Style-Transfer") self.assertFalse(manager.dirty) self.assertEqual( events[-3:], [ ("unload", None), ("load", "Style-Transfer"), ("activate", "Style-Transfer"), ], ) def test_example_index_requires_a_finite_bounded_integer(self): self.assertEqual(parse_example_index(0, 19), 0) self.assertEqual(parse_example_index(18.0, 19), 18) for invalid in (True, "1", None, -1, 19, 1.5, float("nan"), float("inf")): self.assertIsNone(parse_example_index(invalid, 19)) with self.assertRaises(ValueError): parse_example_index(0, -1) def test_fixed_example_manifest_matches_all_enabled_loras(self): tree = ast.parse((ROOT / "app.py").read_text(encoding="utf-8")) assignment = next( node for node in tree.body if isinstance(node, ast.Assign) and any(isinstance(target, ast.Name) and target.id == "EXAMPLES_CONFIG" for target in node.targets) ) examples = ast.literal_eval(assignment.value) self.assertEqual(len(examples), 19) self.assertEqual({example["lora"] for example in examples}, set(ENABLED_ADAPTER_SPECS)) self.assertTrue(all(1 <= len(example["images"]) <= 2 for example in examples)) self.assertEqual(sum(len(example["images"]) == 2 for example in examples), 3) for example in examples: for relative_path in example["images"]: path = (ROOT / relative_path).resolve() self.assertTrue(path.is_relative_to((ROOT / "examples").resolve())) self.assertTrue(path.is_file()) upscaler = next(example for example in examples if example["lora"] == "Upscaler") self.assertEqual(upscaler["prompt"], "Upscale and enhance image detail.") def test_example_endpoint_is_non_gpu_and_adapter_mutation_is_pinned(self): source = (ROOT / "app.py").read_text(encoding="utf-8") tree = ast.parse(source) load_example = next( node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name == "load_example" ) decorators = [ast.unparse(decorator) for decorator in load_example.decorator_list] self.assertEqual(decorators, ["app.api(name='load_example', queue=False)"]) self.assertIn("revision=spec.revision", source) self.assertIn("pipe.unload_lora_weights()", source) self.assertIn("ADAPTER_MANAGER.activate_mode(request.edit_mode)", source) lock_start = source.index("with PIPELINE_LOCK:") activation = source.index("ADAPTER_MANAGER.activate_mode(request.edit_mode)") inference = source.index("result = pipe(", activation) self.assertLess(lock_start, activation) self.assertLess(activation, inference) def test_output_contract_rejects_non_png_bad_dimensions_and_seed(self): with self.assertRaises(ContractError) as caught: encode_png_result(image_bytes("JPEG"), 1) self.assertEqual(caught.exception.code, "I2I_INVALID_OUTPUT") with self.assertRaises(ContractError) as caught: encode_png_result(image_bytes(size=(2048, 64)), 1) self.assertEqual(caught.exception.code, "I2I_INVALID_OUTPUT") with self.assertRaises(ContractError) as caught: encode_png_result(image_bytes(), True) self.assertEqual(caught.exception.code, "I2I_INVALID_OUTPUT") png = serialize_png_candidate(Image.new("RGB", (64, 64))) self.assertTrue(png.startswith(b"\x89PNG\r\n\x1a\n")) if __name__ == "__main__": unittest.main()