# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import numpy as np
import pytest
import torch
from PIL import Image

from vllm.multimodal.parse import (
    AudioProcessorItems,
    ImageProcessorItems,
    MultiModalDataParser,
    VideoProcessorItems,
)

H, W = 480, 640


class AudioMetadataParser(MultiModalDataParser):
    embedding_fields = {
        "audio": {"audio_embeds": "values", "audio_num_tokens": "metadata"},
    }


@pytest.mark.parametrize("allow_missing", [False, True])
def test_audio_metadata_requires_ec_consumer(allow_missing):
    parser = AudioMetadataParser(allow_missing_mm_embeddings=allow_missing)
    data = {"audio_num_tokens": torch.tensor([[3], [5]])}
    if not allow_missing:
        with pytest.raises(ValueError, match="audio_embeds"):
            parser.parse_mm_data({"audio": data})
    else:
        items = parser.parse_mm_data({"audio": data})["audio"]
        assert len(items) == 2
        assert items.get(1)["audio_num_tokens"].item() == 5
        assert items.get_processor_data() == {}


@pytest.mark.parametrize("counts", [[0], [-1], [1.5], [True], [[1, 2]]])
def test_audio_metadata_rejects_invalid_token_counts(counts):
    parser = AudioMetadataParser(allow_missing_mm_embeddings=True)
    with pytest.raises(ValueError, match="positive integer"):
        parser.parse_mm_data({"audio": {"audio_num_tokens": torch.tensor(counts)}})


def test_audio_metadata_checks_supplied_embedding_lengths():
    parser = AudioMetadataParser()
    data = {
        "audio_num_tokens": torch.tensor([3, 5]),
        "audio_embeds": [torch.zeros(3, 8), torch.zeros(5, 8)],
    }
    assert len(parser.parse_mm_data({"audio": data})["audio"]) == 2
    data["audio_num_tokens"] = torch.tensor([3, 4])
    with pytest.raises(ValueError, match="does not match"):
        parser.parse_mm_data({"audio": data})


@pytest.mark.parametrize(
    "image",
    [
        Image.new("RGB", (W, H)),
        # HWC, e.g. from np.array(PIL.Image)
        np.zeros((H, W, 3), dtype=np.uint8),
        torch.zeros((H, W, 3), dtype=torch.uint8),
        # CHW, standard PyTorch / numpy convention
        np.zeros((3, H, W), dtype=np.uint8),
        torch.zeros((3, H, W), dtype=torch.uint8),
    ],
)
def test_image_size_hwc_chw(image):
    """Image sizes must be channel-layout agnostic.

    `get_image_size` determines the multimodal placeholder count; reading an
    HWC array (the layout `np.array(PIL.Image)` produces) as CHW yields a
    bogus size and a placeholder/embedding count mismatch at inference time.
    """
    items = ImageProcessorItems([image])

    assert items.get_image_size(0) == (W, H)


@pytest.mark.parametrize(
    "frame",
    [
        Image.new("RGB", (W, H)),
        np.zeros((H, W, 3), dtype=np.uint8),
        torch.zeros((H, W, 3), dtype=torch.uint8),
        np.zeros((3, H, W), dtype=np.uint8),
        torch.zeros((3, H, W), dtype=torch.uint8),
    ],
)
def test_frame_size_hwc_chw(frame):
    """`get_frame_size` must stay consistent with `get_image_size`."""
    items = VideoProcessorItems([[frame]])

    assert items.get_frame_size(0) == (W, H)


def test_video_with_metadata_tensor_passthrough():
    """Tensor frames pass through unchanged regardless of device: HF video
    processors accept tensors, and device-resident frames (e.g. NVDEC-decoded)
    must not be copied back to host."""
    frames = torch.zeros((4, H, W, 3), dtype=torch.uint8)
    video, metadata = MultiModalDataParser()._get_video_with_metadata(frames)

    assert video is frames
    assert metadata is None


@pytest.mark.skipif(not torch.cuda.is_available(), reason="Requires CUDA")
def test_video_with_metadata_keeps_device_tensor():
    """Device-resident frames (e.g. NVDEC-decoded) pass through as tensors,
    so a device-side HF processor can consume them without a D2H copy."""
    frames = torch.zeros((4, H, W, 3), dtype=torch.uint8, device="cuda")
    video, metadata = MultiModalDataParser()._get_video_with_metadata(frames)

    assert video is frames
    assert metadata is None


@pytest.mark.parametrize(
    "frames",
    [
        [np.zeros((H, W, 3), dtype=np.uint8) for _ in range(2)],
        [torch.zeros((H, W, 3), dtype=torch.uint8) for _ in range(2)],
    ],
)
def test_parse_video_frame_list_as_single_video(frames):
    """A list of decoded frames must represent one video item."""
    items = MultiModalDataParser().parse_mm_data({"video": frames})["video"]

    assert items.get_count() == 1
    video = items.get(0)
    assert isinstance(video, np.ndarray)
    np.testing.assert_array_equal(video, np.stack([np.asarray(f) for f in frames]))


@pytest.mark.parametrize(
    "modality,processor_cls",
    [
        ("audio", AudioProcessorItems),
        ("image", ImageProcessorItems),
        ("video", VideoProcessorItems),
    ],
)
def test_parse_mm_data_accepts_none_cached_item(modality, processor_cls):
    mm_items = MultiModalDataParser().parse_mm_data({modality: [None]})
    items = mm_items[modality]
    assert isinstance(items, processor_cls)
    assert len(items) == 1
    assert items.get(0) is None


def test_cached_audio_items_preserve_positions_during_resampling():
    waveform = np.arange(16, dtype=np.float32)
    parser = MultiModalDataParser(
        target_sr=16000, target_channels=1, audio_resample_method="scipy"
    )
    items = parser.parse_mm_data(
        {"audio": [None, (waveform, 8000), None, (waveform, 16000)]}
    )["audio"]

    assert len(items) == 4
    assert items.get(0) is None
    assert items.get(2) is None
    assert len(items.get(1)) == 32
    np.testing.assert_array_equal(items.get(3), waveform)
