from __future__ import annotations import json import shutil import threading from types import ModuleType from pathlib import Path from adam.assets import AssetRegistry from adam.generations import ( build_generation_plan, generation_output_folder, generation_model_match_score, generation_tools, load_generation_history, parse_chat_generation_request, combine_generation_plans, ) from adam.registry import ToolRegistry from adam.config import ConfigManager from adam.executor import ToolContext from adam.generation_previews import accepts_preview_callback, publish_generation_preview from adam.registry import ToolSpec from adam.tools import ddpm_generator from adam.tools import flow_generator from adam.tools import lora_generator from adam.showcase import build_showcase_plan ROOT = Path(__file__).resolve().parents[1] def test_generation_preview_protocol_keeps_only_latest_preview(tmp_path: Path) -> None: class FakeImage: def save(self, path: Path, *, format: str) -> None: assert format == "PNG" path.write_bytes(b"preview") published: list[dict] = [] running = threading.Event(); running.set() context = ToolContext( root=tmp_path, job_id="PREVIEW1", tool=ToolSpec("ddpm_generator", "DDPM", "test", "Output", "generate_ddpm_images"), cancel_event=threading.Event(), run_event=running, progress_callback=lambda *_args: None, log_callback=lambda *_args: None, preview_callback=published.append, ) output = tmp_path / "output"; output.mkdir() publish_generation_preview( context, output, FakeImage(), image_index=0, image_count=2, step=5, total_steps=20, ) assert (output / ".live_previews" / "PREVIEW1_latest.png").read_bytes() == b"preview" assert published == [{ "path": str(output / ".live_previews" / "PREVIEW1_latest.png"), "epoch": 0, "next_epoch": 0, "prompt": "", "seed": None, "steps": 0, "kind": "generation", "current": 5, "total": 20, "image_index": 1, "image_count": 2, }] def test_preview_callback_protocol_requires_explicit_parameter() -> None: def supported(*, preview_callback): return preview_callback def unsupported(**_settings): return None assert accepts_preview_callback(supported) assert not accepts_preview_callback(unsupported) def test_command_center_parses_quoted_ddpm_generation_request() -> None: parsed = parse_chat_generation_request( 'Generate a "DDPM" image of "Person" for "100" steps, on "DDIM" sampler, with aspect ratio of "16:9"' ) assert parsed is not None assert parsed.provider_hint == "ddpm" assert parsed.prompt == "Person" assert parsed.steps == 100 assert parsed.sampler == "DDIM" assert parsed.aspect_ratio == "16:9" def test_command_center_parses_batch_seed_and_model() -> None: parsed = parse_chat_generation_request( 'Create 3 images of rainy neon streets using model "City Nights" with seed 42 on DDPM sampler at ratio 1:1' ) assert parsed is not None assert parsed.model_query == "City Nights" assert parsed.prompt == "rainy neon streets" assert parsed.image_count == 3 assert parsed.seed == 42 assert parsed.sampler == "DDPM" def test_command_center_does_not_capture_non_image_plans() -> None: assert parse_chat_generation_request("Generate four training previews") is None def test_command_center_subject_matches_completed_model_name() -> None: assert generation_model_match_score("rouge the bat", "Rouge The Bat V2 MADA") > 0 assert generation_model_match_score("rouge the bat", "SpectrogramV3") == 0 def test_command_center_matches_compact_model_names() -> None: assert generation_model_match_score( "neon convenience stores at night", "NeonConvenienceStoresAtNight" ) > 0 def test_command_center_separates_lora_subject_and_positive_prompt() -> None: parsed = parse_chat_generation_request( 'Generate a LoRA image of OrangeCat, Base Model "novaFurryXL", Positive Prompt "OrangeCat, Anthro, female, pretty, sitting on bed, looking at viewer", Negative Prompt "Bad Quality, low effort, missing limbs, poor anatomy"' ) assert parsed is not None assert parsed.provider_hint == "lora" assert parsed.subject == "OrangeCat" assert parsed.base_model_query == "novaFurryXL" assert parsed.prompt == ( "OrangeCat, Anthro, female, pretty, sitting on bed, looking at viewer" ) assert parsed.negative_prompt == "Bad Quality, low effort, missing limbs, poor anatomy" def test_command_center_parses_advanced_generation_overrides() -> None: parsed = parse_chat_generation_request( 'Generate an image, Positive Prompt "city at night", CFG 6.5, ' 'LoRA strength 0.75, denoise strength 0.4, reference strength 70%' ) assert parsed is not None assert parsed.prompt == "city at night" assert parsed.cfg_scale == 6.5 assert parsed.lora_strength == 0.75 assert parsed.denoise_strength == 0.4 assert parsed.reference_strength == 70 def test_command_center_separates_lora_from_base_model_in_natural_phrasing() -> None: parsed = parse_chat_generation_request( 'Generate an image of LoRA OrangeCat, 30 steps, Base Model "waiIllustriousSDXL"' ) assert parsed is not None assert parsed.provider_hint == "lora" assert parsed.model_query == "OrangeCat" assert parsed.base_model_query == "waiIllustriousSDXL" def test_plain_subject_generation_is_not_marked_as_stable_diffusion_prompt() -> None: parsed = parse_chat_generation_request("Generate an image of Minecraft") assert parsed is not None assert parsed.subject == "Minecraft" assert parsed.has_positive_prompt is False assert parsed.provider_hint == "" def test_positive_prompt_is_an_explicit_stable_diffusion_signal() -> None: parsed = parse_chat_generation_request( 'Generate an image, Positive Prompt "Minecraft, blocky world, player"' ) assert parsed is not None assert parsed.prompt == "Minecraft, blocky world, player" assert parsed.has_positive_prompt is True def test_plain_flow_model_suffix_selects_flow_provider() -> None: parsed = parse_chat_generation_request("Generate an image of Minecraft Flow") assert parsed is not None assert parsed.subject == "Minecraft Flow" assert parsed.provider_hint == "flow" def test_external_lora_and_base_model_drop_folders_are_discovered(tmp_path: Path) -> None: lora_folder = tmp_path / "LoRAModelsHere" base_folder = tmp_path / "LoRA StableDiffusionModels Here" lora_folder.mkdir() base_folder.mkdir() (lora_folder / "OrangeCat.safetensors").write_bytes(b"lora") (base_folder / "sdxl-base.safetensors").write_bytes(b"base") assets = AssetRegistry(tmp_path) assets.discover({"tool_folders": {}}) assert any( asset.kind == "model" and asset.trainer == "lora" and asset.name == "OrangeCat" for asset in assets.assets ) assert any( asset.kind == "base_model" and asset.name == "sdxl-base" for asset in assets.assets ) def _registry(tmp_path: Path) -> ToolRegistry: (tmp_path / "config").mkdir() shutil.copy2(ROOT / "config" / "tools.json", tmp_path / "config" / "tools.json") return ToolRegistry(tmp_path) def test_registry_declares_ddpm_image_generator(tmp_path: Path) -> None: registry = _registry(tmp_path) tool = registry.get("ddpm_generator") assert "image_generation" in tool.capabilities assert tool.model_trainers == ("ddpm",) assert [item.id for item in generation_tools(registry)] == [ "ddpm_generator", "flow_generator", "lora_generator", ] def test_registry_declares_lora_image_generator(tmp_path: Path) -> None: tool = _registry(tmp_path).get("lora_generator") assert "image_generation" in tool.capabilities assert "text_prompt" in tool.capabilities assert tool.model_trainers == ("lora",) assert "DPM++ 2M" in tool.generation_options["samplers"] assert "negative_prompt" in tool.arguments assert "base_model_path" in tool.arguments def test_lora_adapter_accepts_cancelled_checkpoints(tmp_path: Path) -> None: cancelled = tmp_path / "Character_cancelled.safetensors" cancelled.write_bytes(b"test") assert lora_generator._lora_file(cancelled) == cancelled def test_registry_declares_flow_image_generator(tmp_path: Path) -> None: tool = _registry(tmp_path).get("flow_generator") assert "image_generation" in tool.capabilities assert tool.model_trainers == ("flow",) assert tool.generation_options["samplers"] == ["Heun", "Euler"] def test_generation_plan_preserves_reproducible_settings(tmp_path: Path) -> None: tool = _registry(tmp_path).get("ddpm_generator") plan = build_generation_plan( tool, model_name="Mario V2", model_path="D:/DDPM/output/Mario", prompt="Explore colorful shapes", image_count=4, steps=75, seed=1234, sampler="DDIM", aspect_ratio="16:9 (Widescreen)", ) assert plan.requires_confirmation is False assert plan.steps[0].tool_id == "ddpm_generator" assert plan.steps[0].arguments == { "model_name": "Mario V2", "model_path": "D:/DDPM/output/Mario", "prompt": "Explore colorful shapes", "image_count": 4, "steps": 75, "seed": 1234, "sampler": "DDIM", "aspect_ratio": "16:9 (Widescreen)", } def test_generation_cycle_keeps_models_in_selected_order(tmp_path: Path) -> None: tool = _registry(tmp_path).get("ddpm_generator") plans = [ build_generation_plan(tool, model_name=name, model_path=f"D:/{name}", prompt="", image_count=3, steps=20, seed=10, sampler="DDIM", aspect_ratio="1:1 (Square)") for name in ("Minecraft", "Roblox") ] cycle = combine_generation_plans(plans, display_seconds=7, show_labels=True) assert cycle.project_name == "Generation Cycle" assert [step.arguments["model_name"] for step in cycle.steps] == ["Minecraft", "Roblox"] assert "7 seconds" in cycle.summary def test_showcase_plan_generates_then_renders_mp4(tmp_path: Path) -> None: registry = _registry(tmp_path) ddpm = registry.get("ddpm_generator") flow = registry.get("flow_generator") plans = [ build_generation_plan( ddpm, model_name="Windows XP", model_path="D:/DDPM/WindowsXP", prompt="", image_count=18, steps=30, seed=100, sampler="DDIM", aspect_ratio="16:9 (Widescreen)", ), build_generation_plan( flow, model_name="Adventure Time", model_path="D:/Flow/AdventureTime", prompt="", image_count=18, steps=30, seed=118, sampler="Heun", aspect_ratio="16:9 (Widescreen)", ), ] settings = [ {"name": "Windows XP", "trainer": "ddpm", "trainer_label": "DDPM", "steps": 30, "sampler": "DDIM", "aspect_ratio": "16:9 (Widescreen)"}, {"name": "Adventure Time", "trainer": "flow", "trainer_label": "Flow Matching", "steps": 30, "sampler": "Heun", "aspect_ratio": "16:9 (Widescreen)"}, ] showcase = build_showcase_plan( plans, title="21 Requests", display_seconds=4, resolution="1080p", model_settings=settings, ) assert showcase.project_name == "Showcase Video" assert [step.tool_id for step in showcase.steps] == [ "ddpm_generator", "flow_generator", "showcase_video_renderer" ] assert showcase.steps[-1].arguments["models"] == settings assert "36 images" in showcase.summary assert "4 seconds" in showcase.summary def test_showcase_rejects_unsupported_image_duration(tmp_path: Path) -> None: tool = _registry(tmp_path).get("ddpm_generator") plan = build_generation_plan( tool, model_name="Test", model_path="D:/Test", prompt="", image_count=12, steps=30, seed=1, sampler="DDIM", aspect_ratio="16:9 (Widescreen)", ) try: build_showcase_plan( [plan], title="Test", display_seconds=2, resolution="720p", model_settings=[{"name": "Test"}], ) except ValueError as exc: assert "3, 4, or 5" in str(exc) else: raise AssertionError("Unsupported showcase duration was accepted") def test_generation_history_reads_images_and_ignores_broken_batches(tmp_path: Path) -> None: good = tmp_path / "data" / "generations" / "good" broken = tmp_path / "data" / "generations" / "broken" good.mkdir(parents=True) broken.mkdir() (good / "image_001.png").write_bytes(b"not decoded by the history loader") (good / "generation.json").write_text( json.dumps( { "provider_id": "ddpm_generator", "provider_name": "DDPM Generator", "model_name": "Mario V2", "model_path": "D:/DDPM/output/Mario", "prompt": "Color study", "seed": 42, "steps": 50, "sampler": "DDIM", "aspect_ratio": "1:1 (Square)", "created_at": "2026-07-31T12:00:00+00:00", } ), encoding="utf-8", ) (broken / "generation.json").write_text("not json", encoding="utf-8") records = load_generation_history(tmp_path) assert len(records) == 1 assert records[0].model_name == "Mario V2" assert records[0].seed == 42 assert records[0].images == (good / "image_001.png",) def test_generation_history_reads_new_model_folder_layout(tmp_path: Path) -> None: folder = generation_output_folder(tmp_path, "ddpm_generator", "Mario V2") image = folder / "20260802_120000_TEST_DDIM_seed_42.png" image.write_bytes(b"image") (folder / "generation_20260802_120000_TEST.json").write_text( json.dumps({"model_name": "Mario V2", "images": [str(image)], "created_at": "2026-08-02T12:00:00+00:00"}), encoding="utf-8", ) records = load_generation_history(tmp_path) assert len(records) == 1 assert records[0].folder == folder assert records[0].images == (image,) def test_ddpm_adapter_writes_images_and_reproducibility_metadata( tmp_path: Path, monkeypatch ) -> None: trainer_root = tmp_path / "connected-ddpm" model = trainer_root / "output" / "Mario" model.mkdir(parents=True) (model / "model_index.json").write_text("{}", encoding="utf-8") script = trainer_root / "appStableDiffusion.py" script.write_text("# test backend", encoding="utf-8") ConfigManager(tmp_path).update( {"tool_folders": {"ddpm_trainer": str(trainer_root)}} ) calls: list[dict] = [] class FakeImage: def save(self, path: Path, *, format: str) -> None: assert format == "PNG" path.write_bytes(b"fake png") backend = ModuleType("fake_ddpm") def generate_images(_model: str, **settings): calls.append(settings) return [FakeImage()] backend.generate_images = generate_images # type: ignore[attr-defined] monkeypatch.setattr(ddpm_generator, "_backend_module", backend) monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) run_event = threading.Event() run_event.set() tool = ToolSpec( id="ddpm_generator", name="DDPM Generator", description="test", category="Output", entry_function="generate_ddpm_images", ) context = ToolContext( root=tmp_path, job_id="TEST1234", tool=tool, cancel_event=threading.Event(), run_event=run_event, progress_callback=lambda _percent, _message: None, log_callback=lambda _message: None, step_delay=0, ) result = ddpm_generator.generate_ddpm_images( context, model_name="Mario", model_path=str(model), prompt="Color study", image_count=2, steps=20, seed=100, sampler="DDIM", aspect_ratio="16:9 (Widescreen)", ) output = Path(str(result["output_folder"])) metadata = json.loads(next(output.glob("generation_*.json")).read_text(encoding="utf-8")) assert metadata["image_seeds"] == [100, 101] assert metadata["prompt_behavior"] == "label_only" assert len(list(output.glob("*.png"))) == 2 assert [call["seed"] for call in calls] == [100, 101] def test_ddpm_adapter_uses_native_reference_image_generation(tmp_path: Path, monkeypatch) -> None: trainer_root = tmp_path / "connected-ddpm" model = trainer_root / "output" / "Mario" model.mkdir(parents=True) (model / "model_index.json").write_text("{}", encoding="utf-8") script = trainer_root / "appStableDiffusion.py" script.write_text("# test backend", encoding="utf-8") reference = tmp_path / "reference.png" reference.write_bytes(b"fake reference") ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) calls: list[dict] = [] class FakeImage: def save(self, path: Path, *, format: str) -> None: path.write_bytes(b"fake png") backend = ModuleType("fake_ddpm_reference") backend.generate_images = lambda *_args, **_kwargs: [FakeImage()] # type: ignore[attr-defined] def generate_reference_images(_model: str, image: str, **settings): calls.append({"image": image, **settings}) return [FakeImage()] backend.generate_reference_images = generate_reference_images # type: ignore[attr-defined] monkeypatch.setattr(ddpm_generator, "_backend_module", backend) monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) event = threading.Event(); event.set() context = ToolContext(tmp_path, "REF1234", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) result = ddpm_generator.generate_ddpm_images( context, "Mario", str(model), "Reference study", 1, 20, 100, "DDIM", "1:1 (Square)", str(reference), 72, ) metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) assert calls == [{"image": str(reference.resolve()), "seed": 100, "num_inference_steps": 20, "batch_size": 1, "sampler": "DDIM", "aspect_ratio": "1:1 (Square)", "reference_strength": 72}] assert metadata["reference_image"] == str(reference.resolve()) assert metadata["reference_strength"] == 72 def test_ddpm_adapter_passes_enabled_custom_dimensions(tmp_path: Path, monkeypatch) -> None: trainer_root = tmp_path / "connected-ddpm" model = trainer_root / "output" / "Mario" model.mkdir(parents=True) (model / "model_index.json").write_text("{}", encoding="utf-8") script = trainer_root / "appStableDiffusion.py" script.write_text("# test backend", encoding="utf-8") ConfigManager(tmp_path).update({"tool_folders": {"ddpm_trainer": str(trainer_root)}}) calls: list[dict] = [] class FakeImage: def save(self, path: Path, *, format: str) -> None: path.write_bytes(b"fake png") backend = ModuleType("fake_ddpm_size") def generate_images(_model: str, **settings): calls.append(settings) return [FakeImage()] backend.generate_images = generate_images # type: ignore[attr-defined] monkeypatch.setattr(ddpm_generator, "_backend_module", backend) monkeypatch.setattr(ddpm_generator, "_backend_script", script.resolve()) monkeypatch.setattr(ddpm_generator.importlib.util, "find_spec", lambda _name: object()) event = threading.Event(); event.set() context = ToolContext(tmp_path, "SIZE1234", ToolSpec("ddpm_generator", "DDPM Generator", "test", "Output", "generate_ddpm_images"), threading.Event(), event, lambda *_args: None, lambda *_args: None, 0) result = ddpm_generator.generate_ddpm_images( context, "Mario", str(model), "Size study", 1, 20, 100, "DDIM", "1:1 (Square)", width=320, height=192, ) metadata = json.loads(next(Path(str(result["output_folder"])).glob("generation_*.json")).read_text(encoding="utf-8")) assert calls[0]["width"] == 320 assert calls[0]["height"] == 192 assert metadata["width"] == 320 assert metadata["height"] == 192 def test_flow_adapter_uses_registered_model_and_tracks_each_seed( tmp_path: Path, monkeypatch ) -> None: flow_root = tmp_path / "connected-flow" model = flow_root / "output_flow_models" / "Rooms" (model / "unet").mkdir(parents=True) (model / "unet" / "config.json").write_text("{}", encoding="utf-8") (model / "flow_model_info.json").write_text( json.dumps({"model_type": "rectified_flow", "model_name": "Rooms"}), encoding="utf-8", ) script = flow_root / "flow_matching_app.py" script.write_text("# test backend", encoding="utf-8") ConfigManager(tmp_path).update( {"tool_folders": {"flow_trainer": str(flow_root)}} ) calls: list[dict] = [] class FakeImage: def save(self, path: Path, *, format: str) -> None: assert format == "PNG" path.write_bytes(b"fake flow png") backend = ModuleType("fake_flow") backend.load_unet = lambda *_args, **_kwargs: object() # type: ignore[attr-defined] def sample_flow(_model, _count, steps, _device, _dtype, seed, method, progress, **settings): progress(steps, steps) calls.append({"seed": seed, "method": method, **settings}) return [FakeImage()] backend.sample_flow = sample_flow # type: ignore[attr-defined] monkeypatch.setattr(flow_generator, "_backend_module", backend) monkeypatch.setattr(flow_generator, "_backend_script", script.resolve()) monkeypatch.setattr(flow_generator, "_loaded_model", None) monkeypatch.setattr(flow_generator, "_loaded_model_path", None) monkeypatch.setattr(flow_generator.importlib.util, "find_spec", lambda _name: object()) run_event = threading.Event() run_event.set() context = ToolContext( root=tmp_path, job_id="FLOW1234", tool=ToolSpec( id="flow_generator", name="Flow Matching Generator", description="test", category="Output", entry_function="generate_flow_images", ), cancel_event=threading.Event(), run_event=run_event, progress_callback=lambda _percent, _message: None, log_callback=lambda _message: None, step_delay=0, ) result = flow_generator.generate_flow_images( context, model_name="Rooms", model_path=str(model), prompt="Room study", image_count=2, steps=8, seed=700, sampler="Heun", aspect_ratio="4:3 (Landscape)", ) output = Path(str(result["output_folder"])) metadata = json.loads(next(output.glob("generation_*.json")).read_text(encoding="utf-8")) assert metadata["provider_id"] == "flow_generator" assert metadata["image_seeds"] == [700, 701] assert [call["seed"] for call in calls] == [700, 701] assert all(call["method"] == "Heun" for call in calls) assert all(call["aspect_ratio"] == "4:3 (Landscape)" for call in calls)