BareGit
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch

import torch
from torch import nn
from safetensors.torch import save_file

from diffusion_cli.config import (
    GenerationDefaults,
    ImageGenerationRequest,
    ModelProfile,
    ModelSource,
    UserConfig,
    TOKENIZER_FILES,
    buildGenerationConfigFromRequest,
    loadUserConfig,
    selectModelProfile,
)
from diffusion_cli.krea2_model import (
    Krea2Conditioning,
    Krea2Config,
    Krea2Model,
)
from diffusion_cli.krea2_text_encoder import (
    KREA2_MAX_TEXT_TOKENS,
    KREA2_PROMPT_PREFIX,
    KREA2_PROMPT_SUFFIX,
    Krea2TextEncoder,
)
from diffusion_cli.krea2_sampling import sampleKrea2, timesteps
from diffusion_cli.quantization import (
    QuantizedLayerSpec,
    QuantizationManifest,
    ScaledFp8Linear,
)
from diffusion_cli.qwen_image_vae import (
    QWEN_IMAGE_LATENT_MEAN,
    QwenImageVae,
    convertWanVaeToDiffusers,
)
from diffusion_cli.model_inspect import inspectModelSource


class Krea2Test(unittest.TestCase):
    def testLongPromptsKeepFixedWindowAndAssistantSuffix(self):
        prefix_ids = list(range(34))
        suffix_ids = list(range(900, 905))

        class FakeTokenizer:
            def __call__(
                self,
                text,
                *,
                return_tensors,
                padding,
                truncation,
                max_length=None,
                add_special_tokens,
            ):
                del return_tensors, truncation, add_special_tokens
                value = text[0] if isinstance(text, list) else text
                if value == KREA2_PROMPT_PREFIX:
                    ids = prefix_ids
                elif value == KREA2_PROMPT_SUFFIX:
                    ids = suffix_ids
                else:
                    ids = prefix_ids + list(range(100, 800))
                    if max_length is not None:
                        ids = ids[:max_length]
                mask = [1] * len(ids)
                if padding == "max_length" and max_length is not None:
                    mask.extend([0] * (max_length - len(ids)))
                    ids.extend([0] * (max_length - len(ids)))
                return {
                    "input_ids": torch.tensor([ids]),
                    "attention_mask": torch.tensor([mask]),
                }

        class FakeModel(nn.Module):
            def forward(self, input_ids, **kwargs):
                del kwargs
                hidden = torch.zeros(
                    input_ids.shape[0],
                    input_ids.shape[1],
                    2560,
                )
                return SimpleNamespace(hidden_states=[hidden] * 36)

        encoder = Krea2TextEncoder(
            None,
            Path("unused"),
            "cpu",
            model=FakeModel(),
            tokenizer=FakeTokenizer(),
        )
        conditioning = encoder.encodePrompt("long prompt")

        self.assertEqual(
            conditioning.hidden_states.shape,
            (1, KREA2_MAX_TEXT_TOKENS, 12, 2560),
        )
        self.assertTrue(conditioning.attention_mask[0, -5:].all())

    def testNamedProfilesCoexistWithLegacyConfiguration(self):
        with tempfile.TemporaryDirectory() as temp_dir:
            root = Path(temp_dir)
            config_path = root / "diffusion.toml"
            config_path.write_text(
                "\n".join(
                    (
                        'default_model = "krea2-turbo"',
                        "",
                        "[models]",
                        f'tokenizer = "{root / "tokenizer"}"',
                        "",
                        "[model_profiles.krea2-turbo]",
                        'architecture = "krea2"',
                        'variant = "turbo"',
                        f'diffusion_model = "{root / "diffusion.safetensors"}"',
                        f'text_encoder = "{root / "text.safetensors"}"',
                        f'vae = "{root / "vae.safetensors"}"',
                        f'tokenizer = "{root / "tokenizer"}"',
                    )
                ),
                encoding="utf-8",
            )
            config = loadUserConfig(config_path)

        self.assertIn("legacy", config.model_profiles)
        self.assertEqual(selectModelProfile(config).name, "krea2-turbo")
        self.assertEqual(config.model_profiles["legacy"].architecture, "z-image")

    def testProfileInspectionReportsCompatibility(self):
        with tempfile.TemporaryDirectory() as temp_dir:
            path = Path(temp_dir) / "diffusion.safetensors"
            save_file(
                {
                    "blocks.0.weight": torch.zeros(1),
                    "txtfusion.projector.weight": torch.zeros(1),
                },
                path,
            )
            summary = inspectModelSource(
                ModelSource(path, "diffusion_model"),
                ModelProfile("krea2-turbo", "krea2", "turbo"),
            )

        self.assertEqual(summary.architecture_guess, "krea2_diffusion")
        self.assertEqual(summary.compatibility, "compatible")

    def testTurboProfileDefaults(self):
        with tempfile.TemporaryDirectory() as temp_dir:
            root = Path(temp_dir)
            tokenizer = root / "tokenizer"
            tokenizer.mkdir()
            for name in TOKENIZER_FILES:
                (tokenizer / name).write_text("{}", encoding="utf-8")
            profile = ModelProfile(
                "krea2-turbo",
                "krea2",
                "turbo",
                diffusion_model=root / "diffusion.safetensors",
                text_encoder=root / "text.safetensors",
                vae=root / "vae.safetensors",
                tokenizer=tokenizer,
            )
            config = UserConfig(
                generation=GenerationDefaults(),
                model_profiles={profile.name: profile},
                default_model=profile.name,
            )
            with patch(
                "diffusion_cli.config.selectDevice",
                return_value=torch.device("cpu"),
            ), patch(
                "diffusion_cli.config.selectDtype",
                return_value=torch.float32,
            ):
                generation = buildGenerationConfigFromRequest(
                    ImageGenerationRequest("a mug"),
                    config,
                )
        self.assertEqual(generation.steps, 8)
        self.assertEqual(generation.cfg, 0.0)
        self.assertEqual(generation.mu, 1.15)

    def testScaledFp8LinearMatchesExplicitDequantization(self):
        layer = ScaledFp8Linear(2, 1, bias=False, dtype=torch.float32)
        stored = torch.tensor([[1.0, 2.0]], dtype=torch.float8_e4m3fn)
        layer.weight.data.copy_(stored)
        layer.weight_scale.data.fill_(2.0)
        value = torch.tensor([[3.0, 4.0]])
        expected = value @ (stored.float() * 2.0).transpose(0, 1)
        self.assertTrue(torch.allclose(layer(value), expected))
        self.assertEqual(layer.weight.dtype, torch.float8_e4m3fn)

    def testReducedModelReturnsFiniteVelocity(self):
        config = Krea2Config(
            features=32,
            timestep_width=8,
            text_width=16,
            heads=4,
            kv_heads=2,
            blocks=1,
            mlp_multiplier=1,
            patch_size=2,
            latent_channels=4,
            text_layers=3,
            text_heads=2,
            text_kv_heads=2,
        )
        model = Krea2Model(config, dtype=torch.float32)
        conditioning = Krea2Conditioning(
            torch.randn(1, 3, 3, 16),
            torch.ones(1, 3, dtype=torch.bool),
        )
        output = model(
            torch.randn(1, 4, 4, 4),
            torch.tensor([0.5]),
            conditioning,
        )
        self.assertEqual(output.shape, (1, 4, 4, 4))
        self.assertTrue(torch.isfinite(output).all())

    def testSamplerSkipsUnconditionalAtZeroCfg(self):
        calls = []

        class FakeModel:
            def __call__(self, latent, timestep, conditioning):
                calls.append(conditioning)
                return torch.ones_like(latent)

        conditioning = Krea2Conditioning(
            torch.zeros(1, 1, 12, 2560),
            torch.ones(1, 1, dtype=torch.bool),
        )
        output = sampleKrea2(
            FakeModel(),
            conditioning,
            batch_size=1,
            height=16,
            width=16,
            seed=7,
            steps=2,
            cfg=0.0,
            device="cpu",
            dtype=torch.float32,
        )
        self.assertEqual(len(calls), 2)
        self.assertEqual(output.shape, (1, 16, 2, 2))

    def testTimestepsDescendAndTurboMuIsStable(self):
        values = timesteps(16, 8, mu=1.15)
        self.assertEqual(len(values), 9)
        self.assertEqual(values[0], 1.0)
        self.assertEqual(values[-1], 0.0)
        self.assertTrue(all(a >= b for a, b in zip(values, values[1:])))

    def testVaeDecodeNormalizesTimeAxisAndRange(self):
        class Result:
            def __init__(self, sample):
                self.sample = sample

        class FakeVae(nn.Module):
            def decode(self, value):
                return Result(
                    torch.zeros(
                        value.shape[0],
                        3,
                        1,
                        value.shape[3] * 8,
                        value.shape[4] * 8,
                    )
                )

        vae = QwenImageVae(None, "cpu", torch.float32, model=FakeVae())
        output = vae.decode(torch.zeros(1, 16, 2, 2))
        self.assertEqual(output.shape, (1, 3, 16, 16))
        self.assertEqual(float(output.mean()), 0.5)

    def testVaeConversionCoversRepresentativeFamilies(self):
        state = {
            "encoder.conv1.weight": torch.zeros(1),
            "decoder.middle.1.proj.weight": torch.zeros(1),
            "encoder.downsamples.0.residual.2.weight": torch.zeros(1),
            "decoder.upsamples.4.shortcut.weight": torch.zeros(1),
            "conv1.weight": torch.zeros(1),
        }
        converted = convertWanVaeToDiffusers(state)
        self.assertIn("encoder.conv_in.weight", converted)
        self.assertIn("decoder.mid_block.attentions.0.proj.weight", converted)
        self.assertIn("encoder.down_blocks.0.conv1.weight", converted)
        self.assertIn(
            "decoder.up_blocks.1.resnets.0.conv_shortcut.weight",
            converted,
        )
        self.assertIn("quant_conv.weight", converted)


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