# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for the multi-token decode flatten in ``TritonMLAImpl.forward_mqa``.

``decode_attention_fwd`` has no causal flag: it launches one program per row of
``q`` and reads that row's KV extent from ``B_Seqlen``. So a ``query_len``-token
decode block is flattened to one row per query token, and intra-block causality
has to be expressed through the per-row sequence lengths -- row ``t`` gets
``seq_len - (query_len - 1) + t``, i.e. the committed prefix plus block tokens
``0..t``. A non-causal draft block instead gives every row the same full extent.

The kernel is replaced by a spy, so these tests assert on the ``block_table``
and ``seq_lens`` the backend actually submits and need no GPU or Triton runtime.
"""

from types import SimpleNamespace
from unittest.mock import patch

import pytest
import torch

from vllm.model_executor.layers.attention.mla_attention import QueryLenSupport
from vllm.v1.attention.backend import AttentionCGSupport
from vllm.v1.attention.backends.mla.triton_mla import (
    TritonMLAImpl,
    TritonMLAMetadataBuilder,
)

NUM_HEADS = 16
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
HEAD_SIZE = KV_LORA_RANK + QK_ROPE_HEAD_DIM
PAGE_SIZE = 16

# Committed context per request, including a fresh request at zero context
# where the whole KV extent is the block itself. A decode block's seq_len
# already counts the block, so seq_len = context + query_len.
CONTEXT_LENS = [0, 1, 37, 512]


def _seq_lens(query_len: int) -> list[int]:
    return [c + query_len for c in CONTEXT_LENS]


def _make_impl(dcp_world_size: int = 1) -> TritonMLAImpl:
    """Only the attributes ``forward_mqa`` reads; __init__ wants a full config."""
    impl = object.__new__(TritonMLAImpl)
    impl.kv_lora_rank = KV_LORA_RANK
    impl.scale = HEAD_SIZE**-0.5
    impl._sm_count = 304
    impl.dcp_world_size = dcp_world_size
    return impl


def _make_metadata(seq_lens: list[int], query_len: int, causal: bool):
    num_decodes = len(seq_lens)
    return SimpleNamespace(
        num_decodes=num_decodes,
        num_decode_tokens=num_decodes * query_len,
        max_seq_len=max(seq_lens),
        causal=causal,
        decode=SimpleNamespace(
            # arange page ids so each request's window is an identifiable slice.
            block_table=torch.arange(num_decodes * 64, dtype=torch.int32).reshape(
                num_decodes, 64
            ),
            seq_lens=torch.tensor(seq_lens, dtype=torch.int32),
        ),
    )


def _run_forward_mqa(
    seq_lens: list[int],
    query_len: int,
    causal: bool,
    q_rows=None,
    dcp_world_size: int = 1,
):
    """Drive forward_mqa with the decode kernel spied out.

    Returns the ``(block_table, seq_lens)`` handed to ``decode_attention_fwd``.
    """
    captured: dict = {}

    def spy(q, kv_cache, kv_c_cache, o, lse, block_table, b_seq_len, *args, **kwargs):
        captured["block_table"] = block_table.detach().clone()
        captured["seq_lens"] = b_seq_len.detach().clone()

    metadata = _make_metadata(seq_lens, query_len, causal)
    rows = metadata.num_decode_tokens if q_rows is None else q_rows
    q = torch.zeros(rows, NUM_HEADS, HEAD_SIZE, dtype=torch.bfloat16)
    kv_cache = torch.zeros(
        len(seq_lens) * 64, PAGE_SIZE, HEAD_SIZE, dtype=torch.bfloat16
    )
    layer = SimpleNamespace(_k_scale=torch.ones(1), layer_name="test")

    with patch("vllm.v1.attention.backends.mla.triton_mla.decode_attention_fwd", spy):
        _make_impl(dcp_world_size).forward_mqa(q, kv_cache, metadata, layer)

    assert captured, "forward_mqa did not reach the decode kernel"
    return captured


def test_multi_token_decode_flags_are_declared_together():
    """Capture probes max_query_len > 1, which only clears the assert in
    build_for_cudagraph_capture once query_len_support raised the threshold."""
    assert (
        TritonMLAMetadataBuilder._cudagraph_support == AttentionCGSupport.UNIFORM_BATCH
    )
    assert TritonMLAMetadataBuilder.query_len_support == QueryLenSupport.UNIFORM


@pytest.mark.parametrize("query_len", [2, 3, 4, 5, 8])
def test_causal_block_rows_see_prefix_plus_own_position(query_len):
    """Causal verify: row t sees seq_len - (query_len - 1) + t KV entries.

    Regression guard: the causal branch used to be skipped entirely, so a verify
    token could attend the draft siblings it is supposed to be checking. Fails
    on unmodified upstream, where no flatten happens at all for causal blocks.
    """
    seq_lens = _seq_lens(query_len)
    captured = _run_forward_mqa(seq_lens, query_len, causal=True)

    got = captured["seq_lens"].tolist()
    want = [s - (query_len - 1) + t for s in seq_lens for t in range(query_len)]
    assert got == want, (
        "causal flatten is not causal: verify row r*query_len+t must get "
        f"seq_len_r - {query_len - 1} + t entries.\n"
        f"  seq_lens {seq_lens}\n  got      {got}\n  expected {want}"
    )

    # Per-request invariants the arithmetic must hold regardless of query_len.
    for r, seq_len in enumerate(seq_lens):
        rows = got[r * query_len : (r + 1) * query_len]
        assert rows == sorted(rows) and len(set(rows)) == query_len, (
            f"request {r} rows must strictly increase, got {rows}"
        )
        assert rows[-1] == seq_len, "last row must see the full sequence"
        assert rows[0] == seq_len - (query_len - 1)
        assert all(n > 0 for n in rows), f"non-positive KV extent in {rows}"


def test_cudagraph_padding_rows_present_no_kv_extent():
    """Causal flattening pins cudagraph padding rows (seq_len 0) to exactly
    zero KV extent, keeping them in the kernel's skipped branch."""
    query_len = 4
    seq_lens = _seq_lens(query_len) + [0]
    captured = _run_forward_mqa(seq_lens, query_len, causal=True)

    padding = captured["seq_lens"].tolist()[-query_len:]
    assert padding == [0] * query_len, (
        f"padding rows must present exactly zero KV extent, got {padding}"
    )


@pytest.mark.parametrize("query_len", [2, 4])
def test_non_causal_block_rows_all_see_the_full_prefix(query_len):
    """The DSpark draft block is generated in one pass, so no row is masked."""
    seq_lens = _seq_lens(query_len)
    captured = _run_forward_mqa(seq_lens, query_len, causal=False)

    want = [s for s in seq_lens for _ in range(query_len)]
    assert captured["seq_lens"].tolist() == want


@pytest.mark.parametrize("causal", [True, False])
def test_block_table_is_expanded_to_match_q_rows(causal):
    """The kernel indexes block_table by program_id, so it needs one row per
    query token -- a short block_table is read out of bounds."""
    query_len = 4
    captured = _run_forward_mqa(_seq_lens(query_len), query_len, causal=causal)

    num_rows = len(CONTEXT_LENS) * query_len
    assert captured["block_table"].shape[0] == num_rows
    assert captured["seq_lens"].shape[0] == num_rows
    # repeat_interleave, not repeat: a request's rows must stay adjacent.
    for r in range(len(CONTEXT_LENS)):
        block = captured["block_table"][r * query_len : (r + 1) * query_len]
        assert torch.equal(block, block[:1].expand_as(block))


@pytest.mark.parametrize("causal", [True, False])
def test_single_token_decode_is_untouched(causal):
    """query_len == 1 must short-circuit: ordinary decode sees no expansion."""
    seq_lens = _seq_lens(1)
    captured = _run_forward_mqa(seq_lens, query_len=1, causal=causal)

    assert captured["seq_lens"].tolist() == seq_lens
    assert captured["block_table"].shape[0] == len(seq_lens)


def test_expansion_factor_follows_the_q_row_count():
    """The kernel indexes block_table and seq_lens by program_id, so the
    expansion must follow q rows, not the token/request counts."""
    num_rows = len(CONTEXT_LENS) * 4
    captured = _run_forward_mqa(_seq_lens(1), query_len=1, causal=True, q_rows=num_rows)

    assert captured["block_table"].shape[0] == num_rows
    assert captured["seq_lens"].shape[0] == num_rows


def test_non_uniform_block_is_rejected():
    """Q rows must divide evenly across the decode requests."""
    with pytest.raises(AssertionError, match="non-uniform decode block"):
        _run_forward_mqa(_seq_lens(1), query_len=1, causal=True, q_rows=6)


def test_causal_multi_token_decode_is_rejected_under_dcp():
    """Per-row extents offset the global seq_len, but DCP passes the rank-local
    slice, so the causal arithmetic would silently address the wrong KV."""
    with pytest.raises(AssertionError, match="not supported with DCP"):
        _run_forward_mqa(_seq_lens(4), query_len=4, causal=True, dcp_world_size=2)

    # Non-causal rows all take the same extent, which stays correct when local.
    _run_forward_mqa(_seq_lens(4), query_len=4, causal=False, dcp_world_size=2)
