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

from copy import deepcopy
from types import SimpleNamespace
from typing import Any

import numpy as np
import pytest
import torch
from PIL import Image

from vllm.config.multimodal import MultiModalConfig
from vllm.model_executor.layers.fusion.mm_input_norm import build_mm_input_norm
from vllm.model_executor.models.qwen2_vl import Qwen2VLProcessingInfo
from vllm.multimodal import MULTIMODAL_REGISTRY
from vllm.multimodal.cache import MultiModalProcessorOnlyCache
from vllm.multimodal.inputs import batched_tensors_equal
from vllm.multimodal.processing.context import InputProcessingContext
from vllm.platforms import current_platform

from ....conftest import ImageTestAssets
from ...utils import build_model_context


@pytest.mark.parametrize(
    ("mm_kwargs", "expected"),
    [
        (
            {"size": {"shortest_edge": 64, "longest_edge": 1024}},
            {
                "size": {"shortest_edge": 64, "longest_edge": 1024},
                "min_pixels": 64,
                "max_pixels": 1024,
            },
        ),
        (
            {"min_pixels": 64, "max_pixels": 1024},
            {
                "min_pixels": 64,
                "max_pixels": 1024,
                "size": {"shortest_edge": 64, "longest_edge": 1024},
            },
        ),
        (
            {
                "size": {"shortest_edge": 64, "longest_edge": 1024},
                "min_pixels": 128,
                "max_pixels": 2048,
            },
            {
                "size": {"shortest_edge": 64, "longest_edge": 1024},
                "min_pixels": 128,
                "max_pixels": 2048,
            },
        ),
        (
            {
                "min_pixels": None,
                "size": {"shortest_edge": 0},
                "images_kwargs": {"max_pixels": 4096},
                "videos_kwargs": {"size": {"longest_edge": 8192}},
            },
            {
                "min_pixels": None,
                "size": {"shortest_edge": 0},
                "images_kwargs": {
                    "max_pixels": 4096,
                    "size": {"longest_edge": 4096},
                },
                "videos_kwargs": {
                    "size": {"longest_edge": 8192},
                    "max_pixels": 8192,
                },
            },
        ),
        (
            {"min_pixels": 0},
            {"min_pixels": 0, "size": {"shortest_edge": 0}},
        ),
        (
            {"size": None, "min_pixels": 64},
            {"size": {"shortest_edge": 64}, "min_pixels": 64},
        ),
    ],
)
def test_complete_mm_processor_size_aliases(
    mm_kwargs: dict[str, object],
    expected: dict[str, object],
) -> None:
    original = deepcopy(mm_kwargs)

    actual = Qwen2VLProcessingInfo._complete_mm_processor_size_aliases(mm_kwargs)

    assert actual == expected
    assert mm_kwargs == original


def test_merge_and_resolve_mm_processor_kwargs_preserves_request_alias_precedence(
    monkeypatch: pytest.MonkeyPatch,
) -> None:
    mm_config = MultiModalConfig(
        mm_processor_kwargs={
            "images_kwargs": {"min_pixels": 64},
        },
        mm_device_do_normalize=False,
    )
    model_config = SimpleNamespace(get_multimodal_config=lambda: mm_config)
    ctx = InputProcessingContext(model_config, tokenizer=None)
    info = Qwen2VLProcessingInfo(ctx)
    monkeypatch.setattr(
        info,
        "get_supported_mm_processor_kwargs",
        lambda: {"images_kwargs": {"size", "min_pixels"}},
    )

    assert info._merge_and_resolve_mm_processor_kwargs(
        {"size": {"shortest_edge": 128}}
    ) == {
        "images_kwargs": {
            "size": {"shortest_edge": 128},
            "min_pixels": 128,
        },
    }


@pytest.mark.parametrize(
    ("mm_kwargs", "default_size", "expected"),
    [
        (
            {},
            {"shortest_edge": 32, "longest_edge": 1024},
            {"shortest_edge": 32, "longest_edge": 1024},
        ),
        (
            {"size": {"shortest_edge": 64}},
            {"shortest_edge": 32, "longest_edge": 1024},
            {"shortest_edge": 64, "longest_edge": 1024},
        ),
        (
            {
                "size": {"shortest_edge": 64, "longest_edge": 2048},
                "min_pixels": 128,
                "max_pixels": 4096,
            },
            {"shortest_edge": 32, "longest_edge": 1024},
            {"shortest_edge": 128, "longest_edge": 4096},
        ),
        (
            {
                "size": {"shortest_edge": 64},
                "max_pixels": 4096,
            },
            None,
            {"shortest_edge": 64, "longest_edge": 4096},
        ),
        (
            {
                "size": {"shortest_edge": 64, "longest_edge": 2048},
                "min_pixels": None,
                "max_pixels": 0,
            },
            {"shortest_edge": 32, "longest_edge": 1024},
            {"shortest_edge": 64, "longest_edge": 0},
        ),
    ],
)
def test_get_vision_size(
    mm_kwargs: dict[str, object],
    default_size: dict[str, object] | None,
    expected: dict[str, object],
) -> None:
    original_mm_kwargs = {
        key: dict(value) if isinstance(value, dict) else value
        for key, value in mm_kwargs.items()
    }
    original_default_size = dict(default_size) if default_size is not None else None

    actual = Qwen2VLProcessingInfo._get_vision_size(
        mm_kwargs,
        default_size=default_size,
    )

    assert actual == expected
    assert mm_kwargs == original_mm_kwargs
    assert default_size == original_default_size


@pytest.mark.parametrize(
    ("modality", "scope"),
    [
        ("image", "images_kwargs"),
        ("video", "videos_kwargs"),
    ],
)
def test_get_vision_info_reads_merged_modality_scope(
    monkeypatch: pytest.MonkeyPatch,
    modality: str,
    scope: str,
) -> None:
    vision_config = SimpleNamespace(
        patch_size=14,
        spatial_merge_size=2,
        temporal_patch_size=2,
    )
    ctx = SimpleNamespace(
        get_hf_config=lambda *_: SimpleNamespace(vision_config=vision_config)
    )
    info = Qwen2VLProcessingInfo(ctx)
    merged = {
        "images_kwargs": {"size": {"shortest_edge": 111}},
        "videos_kwargs": {"size": {"shortest_edge": 222}},
    }
    monkeypatch.setattr(
        info,
        "_merge_and_resolve_mm_processor_kwargs",
        lambda _: merged,
    )

    seen: list[dict[str, object]] = []

    def get_size(
        mm_kwargs: dict[str, object],
        default_size: dict[str, object] | None = None,
    ) -> dict[str, object]:
        seen.append(mm_kwargs)
        return dict(default_size or {})

    monkeypatch.setattr(info, "_get_vision_size", get_size)
    image_processor = SimpleNamespace(size={"shortest_edge": 32, "longest_edge": 1024})

    info._get_vision_info(
        image_width=28,
        image_height=28,
        num_frames=2,
        do_resize=False,
        image_processor=image_processor,
        mm_kwargs={},
        modality=modality,
    )

    assert seen == [merged[scope]]


def test_jina_vl_processing_order() -> None:
    """Jina's document-first prompt keeps cached features and hashes aligned."""
    ctx = build_model_context(
        "jinaai/jina-reranker-m0",
        runner="pooling",
        limit_mm_per_prompt={"image": 2},
        mm_processor_cache_gb=1,
    )
    cache = MultiModalProcessorOnlyCache(ctx.model_config)
    processor = MULTIMODAL_REGISTRY.create_processor(
        ctx.model_config,
        tokenizer=ctx.tokenizer,
    )

    placeholder = "<|vision_start|><|image_pad|><|vision_end|>"
    query_image = Image.new("RGB", (128, 160), color=(255, 0, 0))
    document_image = Image.new("RGB", (192, 128), color=(0, 255, 0))

    def process(images: list[Image.Image]):
        return processor(
            placeholder * len(images),
            mm_items=processor.info.parse_mm_data({"image": images}),
            cache=cache,
        )

    query = process([query_image])
    document = process([document_image])
    pair = process([query_image, document_image])

    pair_items = pair["mm_kwargs"]["image"]
    assert pair["mm_hashes"]["image"] == [
        document["mm_hashes"]["image"][0],
        query["mm_hashes"]["image"][0],
    ]
    assert batched_tensors_equal(
        pair_items[0].get_data(),
        document["mm_kwargs"]["image"][0].get_data(),
    )
    assert batched_tensors_equal(
        pair_items[1].get_data(),
        query["mm_kwargs"]["image"][0].get_data(),
    )
    assert [item.length for item in pair["mm_placeholders"]["image"]] == [
        document["mm_placeholders"]["image"][0].length,
        query["mm_placeholders"]["image"][0].length,
    ]


@pytest.mark.parametrize("model_id", ["Qwen/Qwen2-VL-2B-Instruct"])
@pytest.mark.parametrize(
    ("mm_processor_kwargs", "expected_toks_per_img", "expected_pixels_shape"),
    [
        ({}, 1426, (5704, 1176)),
        ({"min_pixels": 64**2, "max_pixels": 512**2}, 330, (1320, 1176)),
        (
            {
                "size": {
                    "shortest_edge": 64**2,
                    "longest_edge": 512**2,
                },
            },
            330,
            (1320, 1176),
        ),
    ],
)
@pytest.mark.parametrize("num_imgs", [1, 2])
@pytest.mark.parametrize("kwargs_on_init", [True, False])
def test_processor_override(
    image_assets: ImageTestAssets,
    model_id: str,
    mm_processor_kwargs: dict[str, object],
    expected_toks_per_img: int,
    expected_pixels_shape: tuple[int, int],
    num_imgs: int,
    kwargs_on_init: bool,
):
    """Ensure Qwen2VLMultiModalProcessor handles min/max pixels properly."""
    ctx = build_model_context(
        model_id,
        mm_processor_kwargs=mm_processor_kwargs if kwargs_on_init else None,
        limit_mm_per_prompt={"image": num_imgs},
    )
    processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
    tokenizer = processor.info.get_tokenizer()
    hf_processor_mm_kwargs = {} if kwargs_on_init else mm_processor_kwargs

    # Build the image str / prompt based on the number of images we pass
    prompt = "<|vision_start|><|image_pad|><|vision_end|>" * num_imgs
    mm_data = {"image": [image_assets[0].pil_image] * num_imgs}

    processed_inputs = processor(
        prompt,
        mm_items=processor.info.parse_mm_data(mm_data),
        hf_processor_mm_kwargs=hf_processor_mm_kwargs,
    )

    # Ensure we have the right number of placeholders per num_crops size
    hf_processor = processor.info.get_hf_processor(**hf_processor_mm_kwargs)
    image_token_id = tokenizer.convert_tokens_to_ids(hf_processor.image_token)
    img_tok_count = processed_inputs["prompt_token_ids"].count(image_token_id)
    pixel_shape = processed_inputs["mm_kwargs"].get_data()["pixel_values"].shape

    assert img_tok_count == expected_toks_per_img * num_imgs
    assert pixel_shape[0] == expected_pixels_shape[0] * num_imgs
    assert pixel_shape[1] == expected_pixels_shape[1]


@pytest.mark.parametrize("model_id", ["Qwen/Qwen2-VL-2B-Instruct"])
@pytest.mark.parametrize(
    "mm_processor_kwargs",
    [
        {"min_pixels": 28 * 28, "max_pixels": 1280 * 28 * 28},
        {"min_pixels": 28 * 28, "max_pixels": 1283 * 28 * 28},
        {"size": {"shortest_edge": 28 * 28, "longest_edge": 1280 * 28 * 28}},
        {"size": {"shortest_edge": 28 * 28, "longest_edge": 1283 * 28 * 28}},
    ],
)
def test_get_image_size_with_most_features(
    image_assets: ImageTestAssets,
    model_id: str,
    mm_processor_kwargs: dict[str, object],
):
    ctx = build_model_context(
        model_id,
        mm_processor_kwargs=mm_processor_kwargs,
        limit_mm_per_prompt={"image": 1},
    )
    processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)

    hf_processor = processor.info.get_hf_processor(**mm_processor_kwargs)
    merge_size = processor.info.get_hf_config().vision_config.spatial_merge_size

    max_image_size = processor.info.get_image_size_with_most_features()
    max_tokens = processor.info.get_num_image_tokens(
        image_width=max_image_size.width,
        image_height=max_image_size.height,
        image_processor=hf_processor.image_processor,
        mm_kwargs=mm_processor_kwargs,
    )

    prompt = "<|vision_start|><|image_pad|><|vision_end|>"
    for asset in image_assets:
        mm_data = {"image": [asset.pil_image]}
        processed_inputs = processor(
            prompt,
            mm_items=processor.info.parse_mm_data(mm_data),
            hf_processor_mm_kwargs=mm_processor_kwargs,
        )
        grid_thw = processed_inputs["mm_kwargs"].get_data()["image_grid_thw"].tolist()
        t, h, w = grid_thw[0]
        tokens = (t * h * w) // (merge_size**2)
        assert tokens < max_tokens


def _build_qwen2_5_vl_video_mm_data(num_frames: int, fps: float) -> dict[str, Any]:
    video = np.zeros((num_frames, 56, 56, 3), dtype=np.uint8)
    metadata = {
        "fps": fps,
        "duration": num_frames / fps,
        "total_num_frames": num_frames,
        "frames_indices": list(range(num_frames)),
        "video_backend": "opencv",
        "do_sample_frames": False,
    }
    return {"video": [(video, metadata)]}


@pytest.mark.parametrize("model_id", ["Qwen/Qwen2.5-VL-3B-Instruct"])
def test_qwen2_5_vl_video_fps_to_second_per_grid_ts(model_id: str) -> None:
    ctx = build_model_context(model_id, limit_mm_per_prompt={"image": 0, "video": 1})
    processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)
    temporal_patch_size = (
        processor.info.get_hf_processor().video_processor.temporal_patch_size
    )

    prompt = "<|vision_start|><|video_pad|><|vision_end|>"
    for fps in (1.0, 30.0):
        mm_data = _build_qwen2_5_vl_video_mm_data(num_frames=32, fps=fps)
        processed = processor(prompt, mm_items=processor.info.parse_mm_data(mm_data))
        spg = processed["mm_kwargs"]["video"][0].get_data()["second_per_grid_ts"]
        assert float(spg) == pytest.approx(temporal_patch_size / fps, rel=2e-2)


@pytest.mark.usefixtures("default_vllm_config")
@pytest.mark.parametrize(
    "model_id", ["Qwen/Qwen2-VL-2B-Instruct", "Qwen/Qwen2.5-VL-3B-Instruct"]
)
@pytest.mark.parametrize("num_imgs", [1, 2])
def test_mm_device_do_normalize(
    image_assets: ImageTestAssets, model_id: str, num_imgs: int
) -> None:
    """Device-side normalisation must reproduce the on-CPU processor result.

    Runs on any platform: the CPU platform exercises the ``forward_native``
    fallback, accelerators exercise the fused kernel.
    """
    device = current_platform.device_type
    ctx = build_model_context(
        model_id,
        limit_mm_per_prompt={"image": num_imgs},
    )
    ctx.model_config.multimodal_config.mm_device_do_normalize = False
    processor = MULTIMODAL_REGISTRY.create_processor(ctx.model_config)

    # Build the image str / prompt based on the number of images we pass
    prompt = "<|vision_start|><|image_pad|><|vision_end|>" * num_imgs
    mm_data = {"image": [image_assets[0].pil_image] * num_imgs}

    processed_inputs_with_normalize = processor(
        prompt,
        mm_items=processor.info.parse_mm_data(mm_data),
    )
    pixel_values_with_normalize = processed_inputs_with_normalize[
        "mm_kwargs"
    ].get_data()["pixel_values"]
    dtype = pixel_values_with_normalize.dtype

    processed_inputs_without_normalize = processor(
        prompt,
        mm_items=processor.info.parse_mm_data(mm_data),
        hf_processor_mm_kwargs={"do_normalize": False, "do_rescale": False},
    )
    pixel_values_without_normalize = processed_inputs_without_normalize[
        "mm_kwargs"
    ].get_data()["pixel_values"]

    ctx.model_config.multimodal_config.mm_device_do_normalize = True
    input_norm = build_mm_input_norm(ctx.model_config).to(device)

    # With normalisation disabled, the processor emits raw uint8 pixels,
    # matching the production mm_device_do_normalize path.
    assert pixel_values_without_normalize.dtype == torch.uint8
    pixel_values_do_input_norm = input_norm(
        pixel_values_without_normalize.to(device), dtype
    )

    torch.testing.assert_close(
        pixel_values_with_normalize.to(device), pixel_values_do_input_norm
    )
