# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for ApplyRotaryEmb CustomOp dispatch behavior.

This test ensures that RotaryEmbedding classes correctly call the appropriate
ApplyRotaryEmb methods based on the calling context:

1. RotaryEmbedding.forward_native() -> ApplyRotaryEmb.forward_native()
2. RotaryEmbedding.forward_cuda() -> ApplyRotaryEmb.forward() (auto-dispatch)
3. RotaryEmbedding.forward_hip() -> ApplyRotaryEmb.forward() (auto-dispatch)
"""

from dataclasses import dataclass

import pytest
import torch

from vllm.config import (
    CompilationConfig,
    VllmConfig,
    get_cached_compilation_config,
    set_current_vllm_config,
)
from vllm.platforms import current_platform

CUDA_DEVICES = ["cuda:0"]


@dataclass
class RotaryEmbeddingTestCase:
    """Test case configuration for RotaryEmbedding dispatch tests."""

    name: str
    rope_class: type
    rope_kwargs: dict
    method_name: str  # forward_native, forward_cuda, forward
    positions_shape: tuple  # (num_tokens,) or (3, num_tokens) or (4, num_tokens)
    expect_forward_native: bool  # Should call ApplyRotaryEmb.forward_native()
    expect_forward: bool  # Should call ApplyRotaryEmb.forward()


def get_test_cases() -> list[RotaryEmbeddingTestCase]:
    """Generate test cases for all RotaryEmbedding classes."""
    from vllm.model_executor.layers.rotary_embedding.ernie45_vl_rope import (
        Ernie4_5_VLRotaryEmbedding,
    )
    from vllm.model_executor.layers.rotary_embedding.mrope import MRotaryEmbedding

    common_kwargs = {
        "head_size": 128,
        "rotary_dim": 128,
        "max_position_embeddings": 4096,
        "base": 10000,
        "is_neox_style": True,
        "dtype": torch.bfloat16,
    }

    return [
        # MRotaryEmbedding tests
        RotaryEmbeddingTestCase(
            name="MRotaryEmbedding.forward_native",
            rope_class=MRotaryEmbedding,
            rope_kwargs={**common_kwargs, "mrope_section": [16, 24, 24]},
            method_name="forward_native",
            positions_shape=(3, 32),  # 2D for multimodal
            expect_forward_native=True,
            expect_forward=False,
        ),
        RotaryEmbeddingTestCase(
            name="MRotaryEmbedding.forward_cuda_1d",
            rope_class=MRotaryEmbedding,
            rope_kwargs={**common_kwargs, "mrope_section": [16, 24, 24]},
            method_name="forward_cuda",
            positions_shape=(32,),  # 1D triggers apply_rotary_emb path
            expect_forward_native=False,
            expect_forward=True,
        ),
        # Ernie4_5_VLRotaryEmbedding tests
        RotaryEmbeddingTestCase(
            name="Ernie4_5_VLRotaryEmbedding.forward_native",
            rope_class=Ernie4_5_VLRotaryEmbedding,
            rope_kwargs={**common_kwargs, "mrope_section": [22, 22, 20]},
            method_name="forward_native",
            positions_shape=(3, 32),  # 2D for multimodal
            expect_forward_native=True,
            expect_forward=False,
        ),
    ]


def run_dispatch_test(
    test_case: RotaryEmbeddingTestCase,
    device: str,
):
    """Run a dispatch test for a RotaryEmbedding class."""
    vllm_config = VllmConfig(
        compilation_config=CompilationConfig(custom_ops=["all", "+apply_rotary_emb"])
    )
    get_cached_compilation_config.cache_clear()

    with set_current_vllm_config(vllm_config):
        rope = test_case.rope_class(**test_case.rope_kwargs).to(device=device)

        apply_rotary_emb = rope.apply_rotary_emb

        # Verify custom op is enabled
        if test_case.expect_forward_native:
            assert (
                apply_rotary_emb._forward_method != apply_rotary_emb.forward_native
            ), "Test setup error: ApplyRotaryEmb custom op should be enabled"

        # Setup call tracking
        call_tracker = {"forward_native_called": False, "forward_called": False}
        original_forward_native = apply_rotary_emb.forward_native
        original_forward = apply_rotary_emb.forward

        def tracked_forward_native(*args, **kwargs):
            call_tracker["forward_native_called"] = True
            return original_forward_native(*args, **kwargs)

        def tracked_forward(*args, **kwargs):
            call_tracker["forward_called"] = True
            return original_forward(*args, **kwargs)

        apply_rotary_emb.forward_native = tracked_forward_native
        apply_rotary_emb.forward = tracked_forward

        try:
            num_tokens = test_case.positions_shape[-1]
            num_q_heads = 8
            num_kv_heads = 2
            head_size = test_case.rope_kwargs["head_size"]
            max_position = test_case.rope_kwargs["max_position_embeddings"]

            positions = torch.randint(
                0, max_position // 4, test_case.positions_shape, device=device
            )
            query = torch.randn(
                num_tokens, num_q_heads * head_size, dtype=torch.bfloat16, device=device
            )
            key = torch.randn(
                num_tokens,
                num_kv_heads * head_size,
                dtype=torch.bfloat16,
                device=device,
            )

            # Call the method under test
            method = getattr(rope, test_case.method_name)
            method(positions, query.clone(), key.clone())

            # Verify expectations
            if test_case.expect_forward_native:
                assert call_tracker["forward_native_called"], (
                    f"{test_case.name} should call ApplyRotaryEmb.forward_native()"
                )
            if not test_case.expect_forward:
                assert not call_tracker["forward_called"], (
                    f"{test_case.name} should NOT call ApplyRotaryEmb.forward(). "
                    "Bug: when +apply_rotary_emb is enabled, forward_native() "
                    "incorrectly dispatches to CUDA/HIP kernels."
                )
            if test_case.expect_forward:
                assert call_tracker["forward_called"], (
                    f"{test_case.name} should call ApplyRotaryEmb.forward()"
                )
        finally:
            apply_rotary_emb.forward_native = original_forward_native
            apply_rotary_emb.forward = original_forward


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="Skipping CUDA/ROCm only tests."
)
@pytest.mark.parametrize("test_case", get_test_cases(), ids=lambda tc: tc.name)
@pytest.mark.parametrize("device", CUDA_DEVICES)
def test_rotary_embedding_dispatch(
    test_case: RotaryEmbeddingTestCase,
    device: str,
):
    """Test that RotaryEmbedding classes dispatch to the correct ApplyRotaryEmb method.

    - forward_native methods should call ApplyRotaryEmb.forward_native()
    - forward_cuda/forward methods should call ApplyRotaryEmb.forward()
    """
    run_dispatch_test(test_case, device)


@pytest.mark.skipif(not current_platform.is_xpu(), reason="XPU only test.")
@pytest.mark.parametrize(
    "num_tokens,num_heads,head_size,rot_dim,is_neox_style",
    [
        (128, 32, 128, 128, True),  # typical text shape
        (4096, 16, 80, 80, False),  # GPT-J style (interleaved)
        (4096, 32, 128, 64, True),  # partial rotary (rot_dim < head_size)
    ],
    ids=["text", "gptj_interleaved", "partial_rotary"],
)
def test_apply_rotary_emb_xpu_matches_native(
    num_tokens: int,
    num_heads: int,
    head_size: int,
    rot_dim: int,
    is_neox_style: bool,
):
    """ApplyRotaryEmb.forward_xpu (vllm_xpu_kernels SYCL op) must match
    forward_native (the CustomOp default fallback) exactly."""
    from vllm.model_executor.layers.rotary_embedding.common import ApplyRotaryEmb

    vllm_config = VllmConfig(
        compilation_config=CompilationConfig(custom_ops=["all", "+apply_rotary_emb"])
    )
    get_cached_compilation_config.cache_clear()

    with set_current_vllm_config(vllm_config):
        op = ApplyRotaryEmb(enforce_enable=True, is_neox_style=is_neox_style)

        x = torch.randn(
            num_tokens, num_heads, head_size, device="xpu", dtype=torch.bfloat16
        )
        x = x[..., :rot_dim].contiguous()
        cos = torch.randn(num_tokens, rot_dim // 2, device="xpu", dtype=torch.bfloat16)
        sin = torch.randn(num_tokens, rot_dim // 2, device="xpu", dtype=torch.bfloat16)

        expected = op.forward_native(x, cos, sin)
        actual = op.forward_xpu(x, cos, sin)
        torch.testing.assert_close(actual, expected, atol=0, rtol=0)


@pytest.mark.skipif(not current_platform.is_cuda_alike(), reason="CUDA/ROCm only test.")
@pytest.mark.parametrize("num_tokens", [257, 2304])
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
def test_packed_qk_rope_correctness(
    num_tokens: int, dtype: torch.dtype, default_vllm_config
):
    """packed_qk_rope_ must match the per-tensor ApplyRotaryEmb path with
    enable_fp32_compute=True to within one ULP, and must leave the V slice
    untouched.

    Both sides compute in fp32 on identical inputs. The only gap is which
    product the backend keeps exact inside the o1 multiply-add, and that
    follows the register layout, so it is decided by codegen rather than by
    the source and holds neither across backends nor across Triton versions.
    """
    from vllm.model_executor.layers.rotary_embedding.common import ApplyRotaryEmb
    from vllm.model_executor.layers.rotary_embedding.packed_qk_rope import (
        packed_qk_rope_,
    )

    num_heads, head_dim = 12, 128
    rng = torch.Generator(device="cuda").manual_seed(0)

    angles = (
        torch.rand(num_tokens, head_dim // 2, device="cuda", generator=rng) * torch.pi
    )
    freqs_cis = torch.polar(torch.ones_like(angles), angles)
    cos = freqs_cis.real.contiguous()
    sin = freqs_cis.imag.contiguous()

    xqkv = torch.randn(
        num_tokens, 3, num_heads, head_dim, dtype=dtype, device="cuda", generator=rng
    )
    xqkv_ref = xqkv.clone()
    xq, xk, xv = torch.unbind(xqkv_ref, dim=-3)

    op = ApplyRotaryEmb(
        enforce_enable=True, is_neox_style=False, enable_fp32_compute=True
    )
    # op() dispatches to forward_cuda / forward_hip per platform.
    xq_ref = op(xq, cos, sin)
    xk_ref = op(xk, cos, sin)

    packed_qk_rope_(xqkv, freqs_cis)

    # One ULP at the largest magnitude present, not relative: o1 can cancel,
    # so the error scales with the inputs to the sum, not with the result.
    eps = torch.finfo(dtype).eps
    q_atol = eps * xq_ref.abs().max().item()
    k_atol = eps * xk_ref.abs().max().item()
    torch.testing.assert_close(xqkv[:, 0], xq_ref, atol=q_atol, rtol=0)
    torch.testing.assert_close(xqkv[:, 1], xk_ref, atol=k_atol, rtol=0)
    # V is never read or written by the kernel, so it stays exact.
    torch.testing.assert_close(xqkv[:, 2], xv, atol=0, rtol=0)
