# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

"""
End-to-end test for SenseNova-U1 text2img generation.

This test validates that the SenseNova-U1 model generates images that match
expected reference pixel values within a ±10 tolerance.

Equivalent to running:
    python SenseNova-U1/examples/t2i/inference.py \
        --model_path SenseNova/SenseNova-U1-8B-MoT \
        --prompt "Close portrait of an elderly woman ..." \
        --width 1536 --height 2720 \
        --cfg_scale 4.0 --cfg_norm none --timestep_shift 3.0 \
        --num_steps 50 --seed 42 --think
"""

import os
from typing import Any

os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn"

import pytest
from PIL import Image

from tests.helpers.mark import hardware_test
from tests.helpers.runtime import OmniRunner
from vllm_omni.entrypoints.omni import Omni
from vllm_omni.inputs.data import OmniDiffusionSamplingParams

# Reference pixel data extracted from the known-good output image generated by:
#   python examples/offline_inference/sensenova_u1/end2end.py \
#       --prompt "Close portrait of an elderly woman ..." \
#       --width 1536 --height 2720 --seed 42 --num-steps 50 \
#       --cfg-scale 4.0 --timestep-shift 3.0 --cfg-norm none --think
REFERENCE_PIXELS = [
    {"position": (100, 100), "rgb": (247, 249, 250)},
    {"position": (768, 200), "rgb": (176, 135, 97)},
    {"position": (400, 600), "rgb": (193, 188, 180)},
    {"position": (1200, 2000), "rgb": (117, 98, 84)},
    {"position": (750, 500), "rgb": (186, 135, 84)},
    {"position": (300, 1360), "rgb": (202, 157, 107)},
    {"position": (1000, 1800), "rgb": (63, 32, 15)},
    {"position": (500, 2400), "rgb": (135, 116, 104)},
    {"position": (768, 1360), "rgb": (70, 34, 16)},
    {"position": (200, 900), "rgb": (208, 201, 191)},
]

PIXEL_TOLERANCE = 10

DEFAULT_PROMPT = (
    "Close portrait of an elderly woman by a farmhouse window, textured skin, "
    "gentle smile, warm natural light, emotional documentary look. The portrait "
    "should feel polished and natural, with sharp eyes, realistic skin texture, "
    "accurate facial anatomy, and premium lighting that keeps the face as the "
    "main focus."
)

EXPECTED_OUTPUT_SIZE = (1536, 2720)


def _build_sampling_params() -> OmniDiffusionSamplingParams:
    """Build sampling parameters for SenseNova-U1 text2img generation."""
    return OmniDiffusionSamplingParams(
        height=EXPECTED_OUTPUT_SIZE[1],
        width=EXPECTED_OUTPUT_SIZE[0],
        seed=42,
        num_inference_steps=50,
        extra_args={
            "cfg_scale": 4.0,
            "cfg_norm": "none",
            "timestep_shift": 3.0,
            "cfg_interval": (0.0, 1.0),
            "batch_size": 1,
            "think": True,
            "t_eps": 0.02,
        },
    )


def _extract_generated_image(omni_outputs: list) -> Image.Image | None:
    """Extract the generated image from Omni outputs (single-stage DiT)."""
    for req_output in omni_outputs:
        if images := getattr(req_output, "images", None):
            return images[0]
    return None


def _validate_pixels(
    image: Image.Image,
    reference_pixels: list[dict[str, Any]] = REFERENCE_PIXELS,
    tolerance: int = PIXEL_TOLERANCE,
) -> None:
    """Validate that image pixels match expected reference values.

    Args:
        image: The PIL Image to validate.
        reference_pixels: List of dicts with 'position' (x, y) and 'rgb' (R, G, B).
        tolerance: Maximum allowed difference per color channel.

    Raises:
        AssertionError: If any pixel differs beyond tolerance.
    """
    for ref in reference_pixels:
        x, y = ref["position"]
        expected = ref["rgb"]
        actual = image.getpixel((x, y))[:3]
        assert all(abs(a - e) <= tolerance for a, e in zip(actual, expected)), (
            f"Pixel mismatch at ({x}, {y}): expected {expected}, got {actual}"
        )


def _generate_sensenova_u1_image(
    omni: Omni,
    prompt: str = DEFAULT_PROMPT,
) -> Image.Image:
    """Generate an image using SenseNova-U1 model with configured parameters.

    Args:
        omni: The Omni instance to use for generation.
        prompt: The text prompt for image generation.

    Returns:
        The generated PIL Image.

    Raises:
        AssertionError: If no image is generated or size is incorrect.
    """
    sampling_params = _build_sampling_params()

    omni_outputs = list(
        omni.generate(
            prompts={"prompt": prompt, "modalities": ["image"]},
            sampling_params_list=sampling_params,
        )
    )

    generated_image = _extract_generated_image(omni_outputs)
    assert generated_image is not None, "No images generated"
    assert generated_image.size == EXPECTED_OUTPUT_SIZE, f"Expected {EXPECTED_OUTPUT_SIZE}, got {generated_image.size}"

    return generated_image


@pytest.mark.core_model
@pytest.mark.advanced_model
@pytest.mark.diffusion
@hardware_test(res={"cuda": "H100"})
def test_sensenova_u1_text2img(run_level):
    """Test SenseNova-U1 text2img (single-stage diffusion, no deploy YAML)."""
    with OmniRunner(
        "SenseNova/SenseNova-U1-8B-MoT",
        stage_configs_path=None,
    ) as runner:
        generated_image = _generate_sensenova_u1_image(runner.omni)
        if run_level == "advanced_model":
            _validate_pixels(generated_image)


@pytest.mark.core_model
@pytest.mark.diffusion
@pytest.mark.cache
@hardware_test(res={"cuda": "H100"})
def test_sensenova_u1_text2img_cache_dit():
    """Test SenseNova-U1 text2img with Cache-DiT enabled."""
    with OmniRunner(
        "SenseNova/SenseNova-U1-8B-MoT",
        stage_configs_path=None,
        cache_backend="cache_dit",
    ) as runner:
        _generate_sensenova_u1_image(runner.omni)
