Spaces:
Running on Zero
Running on Zero
| 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() | |