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

from types import SimpleNamespace

import numpy as np
import pytest
import torch

from vllm import SamplingParams
from vllm.config.watermarking import WatermarkConfig
from vllm.platforms import current_platform
from vllm.v1.watermarking import (
    DualKeyGumbelWatermarker,
    SupportsSpeculativeDecoding,
    create_watermarker,
    derive_watermark_key,
)
from vllm.v1.watermarking.gpu_sampler import GPUWatermarkSampler
from vllm.v1.watermarking.gumbel import GumbelWatermarker
from vllm.v1.watermarking.spec_decode import (
    DraftWatermarker,
    _resolve_watermark_key,
    create_speculative_draft_watermarker,
    create_speculative_target_watermarker,
    speculative_target_watermark_key,
)
from vllm.v1.watermarking.watermarker import Watermarker, WatermarkSample
from vllm.v1.worker.gpu.buffer_utils import UvaBackedTensor
from vllm.v1.worker.gpu.sample.sampler import Sampler
from vllm.v1.worker.gpu.sample.watermark import (
    draft_watermarking_mask,
    philox_gumbel_sample,
    repeated_context_mask,
)
from vllm.v1.worker.gpu.spec_decode.dspark.speculator import DSparkSpeculator
from vllm.v1.worker.gpu.spec_decode.speculator import DraftModelSpeculator


class StubWatermarker(Watermarker):
    context_width = 1

    def _sample_watermarked(self, logits, contexts):
        return WatermarkSample(torch.tensor([7, 7]), logits + 10)


def _argmax_sampler(logits: torch.Tensor) -> torch.Tensor:
    return logits.argmax(dim=-1)


@pytest.mark.parametrize("algorithm", ["gumbel", "dual_key_gumbel"])
def test_watermarker_contract(algorithm: str):
    watermarker = create_watermarker(
        WatermarkConfig(algorithm=algorithm, key=42, context_width=4)
    )
    logits = torch.zeros(2, 128)
    contexts = torch.tensor([[1, 2, 3, 4], [4, 5, 6, 7]])

    first = watermarker.sample(logits, contexts, _argmax_sampler)
    second = watermarker.sample(logits, contexts, _argmax_sampler)

    assert first.token_ids.shape == (2,)
    assert first.logits.shape == logits.shape
    assert torch.equal(first.token_ids, second.token_ids)


@pytest.mark.parametrize("algorithm", ["gumbel", "dual_key_gumbel"])
@pytest.mark.parametrize(
    "config_overrides",
    [
        {"deduplicate_contexts": "none"},
        {"deduplicate_contexts_max_history": 1023},
    ],
)
def test_gumbel_config_warns_when_context_deduplication_is_weak(
    monkeypatch, algorithm, config_overrides
):
    messages: list[str] = []
    monkeypatch.setattr(
        "vllm.config.watermarking.logger.warning_once",
        lambda message, *, scope: messages.append(message),
    )

    WatermarkConfig(key=42, algorithm=algorithm)
    WatermarkConfig(key=42, algorithm=algorithm, deduplicate_contexts_max_history=1024)
    WatermarkConfig(key=42, algorithm=algorithm, deduplicate_contexts_max_history=None)
    WatermarkConfig(key=42, algorithm=algorithm, **config_overrides)

    assert messages == [
        (
            "Gumbel-max watermarking with context deduplication disabled "
            "or limited to fewer than 1024 positions may increase the frequency of "
            "degenerate generations, including repetition loops. Use "
            "deduplicate_contexts='single_turn' or 'all' with "
            "deduplicate_contexts_max_history at least 1024 or null to mitigate this."
        )
    ]


def test_context_deduplication_history_can_be_unbounded():
    config = WatermarkConfig(key=42, deduplicate_contexts_max_history=None)

    assert config.deduplicate_contexts_max_history is None


def test_large_context_width_warns_but_is_allowed():
    config = WatermarkConfig(key=42, context_width=17)

    with pytest.warns(UserWarning, match="reduce robustness to edits"):
        watermarker = create_watermarker(config)

    assert watermarker.context_width == 17


def test_dual_key_watermarker_uses_domain_separated_keys():
    config = WatermarkConfig(algorithm="dual_key_gumbel", key=42)

    watermarker = create_watermarker(config)
    assert isinstance(watermarker, SupportsSpeculativeDecoding)
    draft = watermarker.draft_watermarker
    target = watermarker.target_watermarker

    assert isinstance(watermarker, DualKeyGumbelWatermarker)
    assert watermarker.alpha == 0.1
    assert draft.prf.key == derive_watermark_key(42, b"key_a")
    assert target.prf.key == derive_watermark_key(42, b"key_b")
    assert target.prf.key != draft.prf.key


def test_dual_key_watermarker_routes_tokens_with_alpha():
    watermarker = DualKeyGumbelWatermarker(key=42, context_width=2, alpha=0.25)
    logits = torch.zeros(2, 16)
    contexts = torch.tensor([[1, 2], [3, 4]])
    key_a = watermarker.draft_watermarker.sample(logits, contexts, lambda _: None)
    key_b = watermarker.target_watermarker.sample(logits, contexts, lambda _: None)

    def route(routing_logits):
        torch.testing.assert_close(
            routing_logits.softmax(dim=-1),
            torch.tensor([[0.75, 0.25], [0.75, 0.25]]),
        )
        return torch.tensor([0, 1])

    sampled = watermarker.sample(logits, contexts, route)

    assert torch.equal(
        sampled.token_ids, torch.stack([key_a.token_ids[0], key_b.token_ids[1]])
    )


def test_dual_key_routing_logits_are_cached():
    alpha = 0.25
    watermarker = DualKeyGumbelWatermarker(key=42, context_width=2, alpha=alpha)
    logits = torch.zeros(2, 16)
    contexts = torch.tensor([[1, 2], [3, 4]])
    recorded: list[torch.Tensor] = []

    def route(routing_logits):
        recorded.append(routing_logits)
        return torch.tensor([0, 1])

    watermarker.sample(logits, contexts, route)
    watermarker.sample(logits, contexts, route)

    expected = torch.tensor([1 - alpha, alpha], dtype=torch.float32).log()
    expected = expected.expand(2, -1)
    assert len(recorded) == 2
    assert all(torch.equal(seen, expected) for seen in recorded)
    assert len(watermarker._routing_logits_cache) == 1


@pytest.mark.parametrize("alpha", [0.0, 0.25, 1.0])
def test_dual_key_watermarker_skips_masked_rows(alpha: float):
    watermarker = DualKeyGumbelWatermarker(key=42, context_width=2, alpha=alpha)
    torch.manual_seed(0)
    logits = torch.randn(4, 32)
    contexts = torch.tensor([[1, 2], [3, 4], [5, 6], [7, 8]])
    skip_mask = torch.tensor([True, False, True, False])

    unmasked = watermarker.sample(logits, contexts, _argmax_sampler)
    masked = watermarker.sample(logits, contexts, _argmax_sampler, skip_mask)

    ordinary = _argmax_sampler(logits)
    assert torch.equal(masked.token_ids[skip_mask], ordinary[skip_mask])
    assert torch.equal(masked.token_ids[~skip_mask], unmasked.token_ids[~skip_mask])
    assert not torch.equal(unmasked.token_ids[skip_mask], ordinary[skip_mask])


def test_dual_key_watermarker_routing_requires_a_random_sampler():
    watermarker = DualKeyGumbelWatermarker(key=42, context_width=2, alpha=0.25)

    with pytest.raises(ValueError, match="random sampler"):
        watermarker.sample(torch.zeros(2, 16), torch.tensor([[1, 2], [3, 4]]))


def test_speculative_decoding_uses_fixed_dual_key_roles():
    watermarker = create_watermarker(
        WatermarkConfig(algorithm="dual_key_gumbel", key=42, alpha=0.25)
    )

    target = create_speculative_target_watermarker(watermarker)
    draft = create_speculative_draft_watermarker(
        watermarker,
        max_num_reqs=1,
        device=torch.device("cpu"),
        allow_target_only=False,
        num_speculative_steps=3,
        deduplicate_contexts="all",
        deduplicate_contexts_max_history=None,
    )

    assert target.prf.key == derive_watermark_key(42, b"key_b")
    assert draft is not None
    assert draft.watermarker.prf.key == derive_watermark_key(42, b"key_a")
    assert draft.num_speculative_steps == 3
    assert draft.deduplicate_contexts == "all"
    assert draft.deduplicate_contexts_max_history is None


def test_recovery_key_rejects_the_unsplit_dual_key_watermarker():
    watermarker = create_watermarker(
        WatermarkConfig(algorithm="dual_key_gumbel", key=42, alpha=0.25)
    )

    with pytest.raises(ValueError, match="keys the target role separately"):
        _resolve_watermark_key(watermarker)

    target = create_speculative_target_watermarker(watermarker)
    assert _resolve_watermark_key(target) == derive_watermark_key(42, b"key_b")


@pytest.mark.parametrize("algorithm", ["gumbel", "dual_key_gumbel"])
def test_config_key_resolution_matches_the_model_runner(algorithm: str):
    config = WatermarkConfig(algorithm=algorithm, key=42, alpha=0.25)
    runtime_key = _resolve_watermark_key(
        create_speculative_target_watermarker(create_watermarker(config))
    )

    assert speculative_target_watermark_key(config) == runtime_key


def test_config_key_resolution_returns_none_when_watermarking_is_disabled():
    assert speculative_target_watermark_key(None) is None


def test_target_only_speculative_watermarking_skips_draft_watermarker():
    watermarker = create_watermarker(WatermarkConfig(algorithm="gumbel", key=42))

    assert not isinstance(watermarker, SupportsSpeculativeDecoding)
    with pytest.raises(ValueError, match="does not support speculative decoding"):
        create_speculative_draft_watermarker(
            watermarker,
            max_num_reqs=1,
            device=torch.device("cpu"),
            allow_target_only=False,
        )
    assert (
        create_speculative_draft_watermarker(
            watermarker,
            max_num_reqs=1,
            device=torch.device("cpu"),
            allow_target_only=True,
        )
        is None
    )


def test_dual_key_derivation_is_stable():
    assert derive_watermark_key(32, b"key_a") == 16368605726115524094
    assert derive_watermark_key(32, b"key_b") == 4799302812959726346


def test_sampling_params_can_disable_watermarking():
    assert SamplingParams().watermarking
    assert not SamplingParams.from_optional(watermarking=False).watermarking


def test_gpu_sampler_warns_when_watermarking_is_enabled_for_greedy(monkeypatch):
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarking = SimpleNamespace(np=np.ones(1, dtype=bool))
    messages: list[str] = []
    monkeypatch.setattr(Sampler, "add_request", lambda *args, **kwargs: None)
    monkeypatch.setattr(
        "vllm.v1.watermarking.gpu_sampler.logger.warning_once", messages.append
    )

    sampler.add_request(0, SamplingParams(temperature=0))
    sampler.add_request(0, SamplingParams(temperature=1))
    sampler.add_request(0, SamplingParams(temperature=0, watermarking=False))

    assert messages == [
        (
            "Watermarking is enabled, but greedy decoding (temperature=0) cannot be "
            "watermarked. This request will use ordinary greedy sampling."
        )
    ]


def test_gpu_sampler_respects_mixed_request_watermarking(monkeypatch):
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = StubWatermarker()
    sampler.deduplicate_contexts = "single_turn"
    sampler.watermarking = SimpleNamespace(
        np=np.array([True, False]), gpu=torch.tensor([True, False])
    )
    sampler.sampling_states = SimpleNamespace(
        temperature=SimpleNamespace(np=np.ones(2), gpu=torch.ones(2)),
        seeds=SimpleNamespace(gpu=torch.zeros(2, dtype=torch.int64)),
    )
    sampler.use_fp64_gumbel = False
    sampler._get_contexts = lambda expanded_idx_mapping: torch.zeros(
        2, 1, dtype=torch.int64
    )
    sampler._get_repeated_contexts = lambda expanded_idx_mapping, contexts: torch.zeros(
        2, dtype=torch.bool
    )
    monkeypatch.setattr(
        "vllm.v1.watermarking.watermarker.gumbel_sample",
        lambda *args, **kwargs: torch.tensor([3, 4]),
    )
    logits = torch.zeros(2, 8)

    sampled, output_logits = sampler._sample_random(
        logits,
        torch.tensor([0, 1]),
        np.array([0, 1]),
        torch.zeros(2, dtype=torch.int64),
        None,
        None,
        False,
    )

    assert torch.equal(sampled, torch.tensor([7, 4]))
    assert torch.equal(output_logits[0], torch.full((8,), 10.0))
    assert torch.equal(output_logits[1], logits[1])


def test_gpu_sampler_skips_watermarking_for_repeated_contexts(monkeypatch):
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = StubWatermarker()
    sampler.deduplicate_contexts = "single_turn"
    sampler.watermarking = SimpleNamespace(
        np=np.array([True, True]), gpu=torch.tensor([True, True])
    )
    sampler.sampling_states = SimpleNamespace(
        temperature=SimpleNamespace(np=np.ones(2), gpu=torch.ones(2)),
        seeds=SimpleNamespace(gpu=torch.zeros(2, dtype=torch.int64)),
    )
    sampler.use_fp64_gumbel = False
    sampler._get_contexts = lambda expanded_idx_mapping: torch.zeros(
        2, 1, dtype=torch.int64
    )
    sampler._get_repeated_contexts = lambda expanded_idx_mapping, contexts: (
        torch.tensor([True, False])
    )
    monkeypatch.setattr(
        "vllm.v1.watermarking.watermarker.gumbel_sample",
        lambda *args, **kwargs: torch.tensor([3, 4]),
    )
    logits = torch.zeros(2, 8)

    sampled, output_logits = sampler._sample_random(
        logits,
        torch.tensor([0, 1]),
        np.array([0, 1]),
        torch.zeros(2, dtype=torch.int64),
        None,
        None,
        False,
    )

    assert torch.equal(sampled, torch.tensor([3, 7]))
    assert torch.equal(output_logits[0], logits[0])
    assert torch.equal(output_logits[1], torch.full((8,), 10.0))


def test_gpu_sampler_can_disable_context_deduplication(monkeypatch):
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = StubWatermarker()
    sampler.deduplicate_contexts = "none"
    sampler.watermarking = SimpleNamespace(
        np=np.array([True, True]), gpu=torch.tensor([True, True])
    )
    sampler.sampling_states = SimpleNamespace(
        temperature=SimpleNamespace(np=np.ones(2), gpu=torch.ones(2)),
        seeds=SimpleNamespace(gpu=torch.zeros(2, dtype=torch.int64)),
    )
    sampler.use_fp64_gumbel = False
    sampler._get_contexts = lambda expanded_idx_mapping: torch.zeros(
        2, 1, dtype=torch.int64
    )
    sampler._get_repeated_contexts = lambda *args: pytest.fail(
        "context deduplication should not run"
    )
    logits = torch.zeros(2, 8)

    sampled, output_logits = sampler._sample_random(
        logits,
        torch.tensor([0, 1]),
        np.array([0, 1]),
        torch.zeros(2, dtype=torch.int64),
        None,
        None,
        False,
    )

    assert torch.equal(sampled, torch.tensor([7, 7]))
    assert torch.equal(output_logits, torch.full((2, 8), 10.0))


def test_repeated_context_mask_ignores_prompt_tokens():
    all_token_ids = torch.tensor(
        [
            [8, 9, 1, 2, 1, 2],
            [3, 4, 1, 2, 3, 4],
            [0, 0, 0, 0, 0, 0],
        ],
        dtype=torch.int32,
    )
    req_indices = torch.tensor([0, 1, -1])
    prompt_lens = torch.tensor([2, 2, 0])
    total_lens = torch.tensor([6, 6, 0])
    contexts = torch.tensor([[1, 2], [3, 4], [-1, -1]])

    repeated = repeated_context_mask(
        all_token_ids,
        req_indices,
        prompt_lens,
        total_lens,
        contexts,
    )

    assert torch.equal(repeated, torch.tensor([True, False, False]))


def test_repeated_context_mask_can_include_prompt_tokens():
    all_token_ids = torch.tensor([[1, 2, 3, 4, 5, 6, 1, 2, 3, 4]])
    req_indices = torch.tensor([0])
    prompt_lens = torch.tensor([4])
    total_lens = torch.tensor([10])
    contexts = torch.tensor([[1, 2, 3, 4]])

    single_turn = repeated_context_mask(
        all_token_ids, req_indices, prompt_lens, total_lens, contexts
    )
    all_history = repeated_context_mask(
        all_token_ids,
        req_indices,
        prompt_lens,
        total_lens,
        contexts,
        include_prompt=True,
    )

    assert not single_turn.item()
    assert all_history.item()


def test_repeated_context_mask_can_skip_partial_contexts():
    all_token_ids = torch.tensor([[10, 11, 12, 13, 1, 2, 3, 4]], dtype=torch.int32)
    req_indices = torch.tensor([0])
    prompt_lens = torch.tensor([4])
    contexts = torch.tensor([[-1, -1, -1, 1]])

    partial = repeated_context_mask(
        all_token_ids,
        req_indices,
        prompt_lens,
        torch.tensor([5]),
        contexts,
        include_prompt=True,
        skip_partial_context=True,
    )
    complete = repeated_context_mask(
        all_token_ids,
        req_indices,
        prompt_lens,
        torch.tensor([8]),
        torch.tensor([[1, 2, 3, 4]]),
        include_prompt=True,
        skip_partial_context=True,
    )

    assert partial.item()
    assert not complete.item()


def test_gpu_sampler_contexts_are_completion_local_for_all_scopes():
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = SimpleNamespace(context_width=4)
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(gpu=torch.tensor([[10, 11, 12, 13, 1, 2]])),
        prompt_len=SimpleNamespace(gpu=torch.tensor([4])),
        total_len=SimpleNamespace(gpu=torch.tensor([6])),
    )
    request_indices = torch.tensor([0])

    sampler.deduplicate_contexts = "single_turn"
    single_turn = sampler._get_contexts(request_indices)
    sampler.deduplicate_contexts = "all"
    all_history = sampler._get_contexts(request_indices)

    assert torch.equal(single_turn, torch.tensor([[-1, -1, 1, 2]]))
    assert torch.equal(all_history, single_turn)


def test_gpu_sampler_all_history_skips_partial_context():
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = SimpleNamespace(context_width=4)
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(gpu=torch.tensor([[10, 11, 12, 13, 1, 2]])),
        prompt_len=SimpleNamespace(gpu=torch.tensor([4])),
        total_len=SimpleNamespace(gpu=torch.tensor([6])),
    )
    sampler.deduplicate_contexts_max_history = None
    request_indices = torch.tensor([0])
    contexts = sampler._get_contexts(request_indices)

    sampler.deduplicate_contexts = "single_turn"
    single_turn = sampler._get_repeated_contexts(request_indices, contexts)
    sampler.deduplicate_contexts = "all"
    all_history = sampler._get_repeated_contexts(request_indices, contexts)

    assert not single_turn.item()
    assert all_history.item()


def test_repeated_context_mask_respects_max_history():
    all_token_ids = torch.tensor([[1, 2, 3, 4, 5, 6, 1, 2, 3, 4]])
    req_indices = torch.tensor([0])
    prompt_lens = torch.tensor([0])
    total_lens = torch.tensor([10])
    contexts = torch.tensor([[1, 2, 3, 4]])

    full_history = repeated_context_mask(
        all_token_ids, req_indices, prompt_lens, total_lens, contexts
    )
    last_six = repeated_context_mask(
        all_token_ids, req_indices, prompt_lens, total_lens, contexts, max_history=6
    )
    last_five = repeated_context_mask(
        all_token_ids, req_indices, prompt_lens, total_lens, contexts, max_history=5
    )

    assert full_history.item()
    assert last_six.item()
    assert not last_five.item()


@pytest.mark.parametrize("max_history", [0, -1])
def test_repeated_context_mask_rejects_non_positive_max_history(max_history: int):
    with pytest.raises(ValueError, match="max_history must be positive or None"):
        repeated_context_mask(
            torch.tensor([[1, 2, 3, 4]], dtype=torch.int32),
            torch.tensor([0], dtype=torch.int32),
            torch.tensor([0], dtype=torch.int32),
            torch.tensor([4], dtype=torch.int32),
            torch.tensor([[1, 2, 3, 4]], dtype=torch.int32),
            max_history=max_history,
        )


def test_gpu_sampler_wires_request_history_scope():
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(
            gpu=torch.tensor([[1, 2, 3, 4, 5, 6, 1, 2, 3, 4]])
        ),
        prompt_len=SimpleNamespace(gpu=torch.tensor([4])),
        total_len=SimpleNamespace(gpu=torch.tensor([10])),
    )
    sampler.deduplicate_contexts_max_history = None
    request_indices = torch.tensor([0])
    contexts = torch.tensor([[1, 2, 3, 4]])

    sampler.deduplicate_contexts = "single_turn"
    single_turn = sampler._get_repeated_contexts(request_indices, contexts)
    sampler.deduplicate_contexts = "all"
    all_history = sampler._get_repeated_contexts(request_indices, contexts)

    assert not single_turn.item()
    assert all_history.item()


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
@pytest.mark.parametrize("max_history", [5, 6])
@pytest.mark.parametrize("include_prompt", [False, True])
def test_repeated_context_mask_max_history_accelerator_parity(
    max_history: int, include_prompt: bool
):
    inputs = (
        torch.tensor([[1, 2, 3, 4, 5, 6, 1, 2, 3, 4]]),
        torch.tensor([0]),
        torch.tensor([0]),
        torch.tensor([10]),
        torch.tensor([[1, 2, 3, 4]]),
    )

    expected = repeated_context_mask(
        *inputs, max_history=max_history, include_prompt=include_prompt
    )
    actual = repeated_context_mask(
        *(value.cuda() for value in inputs),
        max_history=max_history,
        include_prompt=include_prompt,
    ).cpu()

    assert torch.equal(actual, expected)


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
def test_repeated_context_mask_partial_context_accelerator_parity():
    inputs = (
        torch.tensor([[10, 11, 12, 13, 1]], dtype=torch.int32),
        torch.tensor([0], dtype=torch.int32),
        torch.tensor([4], dtype=torch.int32),
        torch.tensor([5], dtype=torch.int32),
        torch.tensor([[-1, -1, -1, 1]], dtype=torch.int64),
    )

    expected = repeated_context_mask(
        *inputs, include_prompt=True, skip_partial_context=True
    )
    actual = repeated_context_mask(
        *(value.cuda() for value in inputs),
        include_prompt=True,
        skip_partial_context=True,
    ).cpu()

    assert torch.equal(actual, expected)


def _unaligned_max_history_inputs() -> tuple[torch.Tensor, ...]:
    """One row whose only repeated context ends at position 680.

    680 is not a multiple of the kernel's 512-position block, so the scan window
    selected by ``max_history`` starts in the middle of a block.
    """
    history = list(range(1_000, 2_200))
    history[676:680] = [1, 2, 3, 4]
    return (
        torch.tensor([history], dtype=torch.int32),
        torch.tensor([0], dtype=torch.int32),
        torch.tensor([0], dtype=torch.int32),
        torch.tensor([len(history)], dtype=torch.int32),
        torch.tensor([[1, 2, 3, 4]], dtype=torch.int32),
    )


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
@pytest.mark.parametrize(
    "max_history,expected", [(1200 - 680, True), (1200 - 680 - 1, False)]
)
def test_repeated_context_mask_unaligned_max_history_accelerator_parity(
    max_history: int, expected: bool
):
    inputs = _unaligned_max_history_inputs()

    reference = repeated_context_mask(*inputs, max_history=max_history)
    actual = repeated_context_mask(
        *(value.cuda() for value in inputs), max_history=max_history
    ).cpu()

    assert torch.equal(reference, torch.tensor([expected]))
    assert torch.equal(actual, reference)


def _repeated_context_inputs(
    context_width: int, device: str = "cpu"
) -> tuple[torch.Tensor, ...]:
    """Build repeated, unique, and padded request rows for mask tests."""
    prompt = [101, 102, 103, 104]
    repeated_output = [*range(1, context_width + 1)] * 2
    unique_output = [*range(1, context_width + 2)]
    max_len = len(prompt) + len(repeated_output)
    rows = [
        prompt + repeated_output,
        prompt + unique_output,
        [],
    ]
    all_token_ids = torch.zeros((3, max_len), dtype=torch.int32, device=device)
    for row, token_ids in enumerate(rows):
        all_token_ids[row, : len(token_ids)] = torch.tensor(
            token_ids, dtype=torch.int32, device=device
        )
    return (
        all_token_ids,
        torch.tensor([0, 1, -1], dtype=torch.int32, device=device),
        torch.tensor([len(prompt), len(prompt), 0], dtype=torch.int32, device=device),
        torch.tensor([len(rows[0]), len(rows[1]), 0], dtype=torch.int32, device=device),
        torch.tensor(
            [
                repeated_output[-context_width:],
                unique_output[-context_width:],
                [-1] * context_width,
            ],
            dtype=torch.int64,
            device=device,
        ),
    )


@pytest.mark.parametrize("context_width", [1, 3, 4, 16, 17])
def test_repeated_context_mask(context_width: int):
    repeated = repeated_context_mask(*_repeated_context_inputs(context_width))

    assert torch.equal(repeated, torch.tensor([True, False, False]))


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
@pytest.mark.parametrize("context_width", [1, 3, 4, 16, 17])
def test_repeated_context_mask_accelerator_parity(context_width: int):
    cpu_inputs = _repeated_context_inputs(context_width)
    accelerator_inputs = list(_repeated_context_inputs(context_width, "cuda"))
    contexts = accelerator_inputs[-1]
    storage = torch.empty(
        contexts.shape[0], contexts.shape[1] * 2, dtype=contexts.dtype, device="cuda"
    )
    storage[:, ::2] = contexts
    accelerator_inputs[-1] = storage[:, ::2]

    expected = repeated_context_mask(*cpu_inputs)
    actual = repeated_context_mask(*accelerator_inputs).cpu()

    assert torch.equal(actual, expected)


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
def test_repeated_context_mask_scans_multiple_blocks():
    context = [1, 2, 3, 4]
    repeated_output = [*range(10_000, 11_030), *context, 99, *context]
    unique_output = list(range(20_000, 21_039))
    prompt = [101, 102, 103, 104]
    rows = [prompt + repeated_output, prompt + unique_output]
    all_token_ids = torch.zeros((2, len(rows[0])), dtype=torch.int32)
    for row, token_ids in enumerate(rows):
        all_token_ids[row, : len(token_ids)] = torch.tensor(token_ids)
    inputs = (
        all_token_ids,
        torch.tensor([0, 1], dtype=torch.int32),
        torch.tensor([len(prompt), len(prompt)], dtype=torch.int32),
        torch.tensor([len(rows[0]), len(rows[1])], dtype=torch.int32),
        torch.tensor([context, unique_output[-4:]], dtype=torch.int64),
    )

    expected = repeated_context_mask(*inputs)
    actual = repeated_context_mask(*(value.cuda() for value in inputs)).cpu()

    assert torch.equal(expected, torch.tensor([True, False]))
    assert torch.equal(actual, expected)


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
def test_repeated_context_mask_last_request_at_capacity():
    max_model_len = 1024
    all_token_ids = torch.arange(
        2 * max_model_len, dtype=torch.int32, device="cuda"
    ).reshape(2, max_model_len)
    inputs = (
        all_token_ids,
        torch.tensor([1], dtype=torch.int32, device="cuda"),
        torch.tensor([0, 0], dtype=torch.int32, device="cuda"),
        torch.tensor([max_model_len, max_model_len], dtype=torch.int32, device="cuda"),
        all_token_ids[1, -4:].to(torch.int64).unsqueeze(0),
    )

    repeated = repeated_context_mask(*inputs)

    assert not repeated.item()


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
def test_gpu_sampler_uses_fused_gumbel_for_repeated_contexts():
    device = torch.device("cuda")
    logits = torch.tensor([[0.0, 1.0, 2.0, 3.0], [3.0, 2.0, 1.0, 0.0]], device=device)
    all_token_ids = torch.tensor(
        [[101, 102, 1, 2, 1, 2], [101, 102, 3, 4, 5, 6]],
        dtype=torch.int32,
        device=device,
    )
    request_indices = torch.tensor([0, 1], dtype=torch.int64, device=device)
    request_indices_np = np.array([0, 1])
    temperatures = torch.ones(2, device=device)
    seeds = torch.tensor([11, 22], dtype=torch.int64, device=device)
    positions = torch.tensor([6, 6], dtype=torch.int64, device=device)

    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = GumbelWatermarker(key=42, context_width=1)
    sampler.deduplicate_contexts = "single_turn"
    sampler.deduplicate_contexts_max_history = None
    sampler.watermarking = SimpleNamespace(
        np=np.ones(2, dtype=bool), gpu=torch.ones(2, dtype=torch.bool, device=device)
    )
    sampler.sampling_states = SimpleNamespace(
        temperature=SimpleNamespace(np=np.ones(2), gpu=temperatures),
        seeds=SimpleNamespace(gpu=seeds),
    )
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(gpu=all_token_ids),
        prompt_len=SimpleNamespace(
            gpu=torch.tensor([2, 2], dtype=torch.int32, device=device)
        ),
        total_len=SimpleNamespace(
            gpu=torch.tensor([6, 6], dtype=torch.int32, device=device)
        ),
    )
    sampler.use_fp64_gumbel = False

    contexts = sampler._get_contexts(request_indices)
    repeated = sampler._get_repeated_contexts(request_indices, contexts)
    expected = philox_gumbel_sample(
        logits,
        contexts,
        42,
        skip_mask=repeated,
        expanded_idx_mapping=request_indices,
        temperatures=temperatures,
        seeds=seeds,
        positions=positions,
    )

    actual, output_logits = sampler._sample_random(
        logits,
        request_indices,
        request_indices_np,
        positions,
        None,
        None,
        False,
    )

    assert torch.equal(repeated, torch.tensor([True, False], device=device))
    assert torch.equal(actual, expected)
    assert output_logits is logits


def test_gpu_sampler_skips_watermarking_for_greedy_batch(monkeypatch):
    class StubWatermarker:
        context_width = 1

        def sample(self, logits, contexts, random_sampler=None, skip_mask=None):
            raise AssertionError("watermarker should not run for greedy requests")

    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = StubWatermarker()
    sampler.watermarking = SimpleNamespace(
        np=np.array([True, True]), gpu=torch.tensor([True, True])
    )
    sampler.sampling_states = SimpleNamespace(
        temperature=SimpleNamespace(np=np.zeros(2), gpu=torch.zeros(2)),
    )
    expected = (torch.tensor([3, 4]), torch.zeros(2, 8))
    monkeypatch.setattr(
        "vllm.v1.watermarking.gpu_sampler.Sampler._sample_random",
        lambda *args, **kwargs: expected,
    )

    actual = sampler._sample_random(
        torch.zeros(2, 8),
        torch.tensor([0, 1]),
        np.array([0, 1]),
        torch.zeros(2, dtype=torch.int64),
        None,
        None,
        False,
    )

    assert actual is expected


def test_gpu_sampler_builds_speculative_contexts_from_drafts():
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = SimpleNamespace(context_width=2)
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(gpu=torch.tensor([[100, 101, 10, 11, 0, 0]])),
        prompt_len=SimpleNamespace(gpu=torch.tensor([2])),
        total_len=SimpleNamespace(gpu=torch.tensor([4])),
    )

    contexts = sampler._get_contexts(
        torch.tensor([0, 0, 0]),
        torch.tensor([0, 1, 2]),
        torch.tensor([11, 20, 21]),
    )

    assert torch.equal(contexts, torch.tensor([[10, 11], [11, 20], [20, 21]]))


def test_gpu_sampler_deduplicates_contexts_within_speculative_block():
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = SimpleNamespace(context_width=2)
    sampler.deduplicate_contexts = "single_turn"
    sampler.deduplicate_contexts_max_history = None
    sampler.num_speculative_tokens = 4
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(gpu=torch.tensor([[100, 1, 2, 0, 0, 0]])),
        prompt_len=SimpleNamespace(gpu=torch.tensor([1])),
        total_len=SimpleNamespace(gpu=torch.tensor([3])),
    )
    request_indices = torch.tensor([0, 0, 0, 0, 0])
    local_positions = torch.tensor([0, 1, 2, 3, 4])
    contexts = sampler._get_contexts(
        request_indices,
        local_positions,
        torch.tensor([0, 3, 4, 1, 2]),
    )

    repeated = sampler._get_repeated_contexts(
        request_indices, contexts, local_positions
    )

    assert torch.equal(
        contexts,
        torch.tensor([[1, 2], [2, 3], [3, 4], [4, 1], [1, 2]]),
    )
    assert torch.equal(repeated, torch.tensor([False, False, False, False, True]))

    sampler.deduplicate_contexts_max_history = 3
    repeated = sampler._get_repeated_contexts(
        request_indices, contexts, local_positions
    )
    assert not repeated.any()


def test_gpu_sampler_all_scope_skips_only_partial_speculative_contexts():
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = SimpleNamespace(context_width=2)
    sampler.deduplicate_contexts = "all"
    sampler.deduplicate_contexts_max_history = None
    sampler.num_speculative_tokens = 3
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(gpu=torch.tensor([[100, 101, 10, 0, 0]])),
        prompt_len=SimpleNamespace(gpu=torch.tensor([2])),
        total_len=SimpleNamespace(gpu=torch.tensor([3])),
    )
    request_indices = torch.tensor([0, 0, 0])
    local_positions = torch.tensor([0, 1, 2])
    contexts = sampler._get_contexts(
        request_indices,
        local_positions,
        torch.tensor([10, 20, 30]),
    )

    repeated = sampler._get_repeated_contexts(
        request_indices, contexts, local_positions
    )

    assert torch.equal(contexts, torch.tensor([[-1, 10], [10, 20], [20, 30]]))
    assert torch.equal(repeated, torch.tensor([True, False, False]))


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
@pytest.mark.parametrize("max_history,expected_last", [(None, True), (3, False)])
def test_speculative_context_repetition_accelerator_parity(
    max_history: int | None, expected_last: bool
):
    all_token_ids = torch.tensor([[10, 11]])
    req_indices = torch.tensor([0, 0, 0, 0, 0])
    prompt_lens = torch.tensor([0])
    total_lens = torch.tensor([2])
    local_positions = torch.tensor([0, 1, 2, 3, 4])
    contexts = torch.tensor([[1, 2], [2, 3], [3, 4], [4, 1], [1, 2]])

    expected = repeated_context_mask(
        all_token_ids,
        req_indices,
        prompt_lens,
        total_lens,
        contexts,
        history_offsets=local_positions,
        local_positions=local_positions,
        num_speculative_steps=4,
        max_history=max_history,
    )
    actual = repeated_context_mask(
        all_token_ids.cuda(),
        req_indices.cuda(),
        prompt_lens.cuda(),
        total_lens.cuda(),
        contexts.cuda(),
        history_offsets=local_positions.cuda(),
        local_positions=local_positions.cuda(),
        num_speculative_steps=4,
        max_history=max_history,
    ).cpu()

    assert torch.equal(
        expected, torch.tensor([False, False, False, False, expected_last])
    )
    assert torch.equal(actual, expected)


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
@pytest.mark.parametrize("steps_layout", ["per_row", "broadcast"])
@pytest.mark.parametrize("req_indices_layout", ["contiguous", "column_view"])
def test_draft_context_deduplication_accelerator_parity(
    steps_layout, req_indices_layout
):
    """Parity with the CPU reference for the layouts the speculators pass.

    A draft step is one 0-d tensor broadcast over the batch, and DSpark samples
    a step at a time, so its request indices are an `idx_map[:, i]` column. Both
    are read per row, so a row that ignores the layout consults another row's
    step or another request's history.
    """
    num_rows, num_steps, step = 4, 3, 1
    all_token_ids = torch.tensor([[10, 11, 20, 11], [30, 31, 11, 20]])
    prompt_lens = torch.tensor([0, 0])
    total_lens = torch.tensor([4, 4])
    contexts = torch.tensor([[40, 41], [31, 11], [10, 11], [50, 51]])
    prior_contexts = torch.zeros(num_rows, num_steps, 2, dtype=torch.int64)
    prior_contexts[0, 0] = torch.tensor([40, 41])
    enabled = torch.ones(num_rows, dtype=torch.bool)
    req_indices = torch.tensor([0, 1, 0, 1])

    def run(all_token_ids, req_indices, steps, prior):
        return draft_watermarking_mask(
            all_token_ids,
            req_indices,
            prompt_lens.to(all_token_ids.device),
            total_lens.to(all_token_ids.device),
            prior,
            contexts.to(all_token_ids.device),
            steps,
            enabled.to(all_token_ids.device),
            max_history=3,
            include_prompt=False,
        )

    reference_prior = prior_contexts.clone()
    reference = run(
        all_token_ids,
        req_indices,
        torch.full((num_rows,), step),
        reference_prior,
    )

    # Row 0 repeats its own step-0 context, rows 1 and 2 repeat a committed
    # context of their own request, row 3 repeats nothing.
    assert torch.equal(reference, torch.tensor([False, False, False, True]))
    expected_prior = prior_contexts.clone()
    expected_prior[:, step] = contexts
    assert torch.equal(reference_prior, expected_prior)

    if steps_layout == "per_row":
        steps = torch.full((num_rows,), step).cuda()
    else:
        steps = torch.tensor(step).cuda().expand(num_rows)
        assert steps.stride(0) == 0

    if req_indices_layout == "contiguous":
        gpu_req_indices = req_indices.cuda()
    else:
        gpu_req_indices = (
            req_indices.repeat_interleave(num_steps).view(num_rows, num_steps).cuda()
        )[:, 1]
        assert gpu_req_indices.stride(0) == num_steps

    actual_prior = prior_contexts.cuda()
    actual = run(all_token_ids.cuda(), gpu_req_indices, steps, actual_prior).cpu()

    assert torch.equal(actual, reference)
    assert torch.equal(actual_prior.cpu(), expected_prior)


def test_gpu_sampler_builds_chunked_multi_request_speculative_contexts():
    sampler = object.__new__(GPUWatermarkSampler)
    sampler.watermarker = SimpleNamespace(context_width=2)
    sampler.req_states = SimpleNamespace(
        all_token_ids=SimpleNamespace(
            gpu=torch.tensor(
                [
                    [100, 101, 10, 11, 0, 0],
                    [200, 201, 30, 31, 0, 0],
                    [300, 301, 50, 51, 0, 0],
                ]
            )
        ),
        prompt_len=SimpleNamespace(gpu=torch.tensor([2, 2, 2])),
        total_len=SimpleNamespace(gpu=torch.tensor([4, 4, 4])),
    )

    contexts = sampler._get_contexts(
        torch.tensor([2, 2, 1, 1]),
        torch.tensor([0, 1, 0, 1]),
        torch.tensor([51, 60, 31, 40]),
    )

    assert torch.equal(
        contexts,
        torch.tensor([[50, 51], [51, 60], [30, 31], [31, 40]]),
    )


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
def test_draft_context_deduplication_broadcasts_the_draft_step():
    """The draft step is one 0-d tensor broadcast over the batch.

    The context occurs four positions back, which `max_history - step` puts
    just outside the scanned window, so a row that mistakes another row's step
    for its own reports the opposite verdict. The neighbours leg is the
    deterministic control: a row that reads a step of 0 instead of 2 widens its
    window by two positions and finds the context. The production leg asserts
    the layout the speculators pass, where the stride-less read lands outside
    the allocation and can return the right value by chance.
    """
    num_rows = 4
    all_token_ids = torch.tensor([[3, 4, 9, 9, 9, 9], [3, 4, 9, 9, 9, 9]])
    req_indices = torch.tensor([0, 0, 1, 1])
    prompt_lens = torch.tensor([0, 0])
    total_lens = torch.tensor([6, 6])
    prior_contexts = torch.zeros(num_rows, 3, 2, dtype=torch.int64)
    contexts = torch.tensor([[3, 4]]).expand(num_rows, 2).contiguous()
    enabled = torch.ones(num_rows, dtype=torch.bool)

    def run(steps, device):
        prior = prior_contexts.clone().to(device)
        active = draft_watermarking_mask(
            all_token_ids.to(device),
            req_indices.to(device),
            prompt_lens.to(device),
            total_lens.to(device),
            prior,
            contexts.to(device),
            steps,
            enabled.to(device),
            max_history=4,
            include_prompt=False,
        )
        return active.cpu(), prior.cpu()

    expected, expected_prior = run(torch.full((num_rows,), 2), torch.device("cpu"))
    assert torch.equal(expected, torch.ones(num_rows, dtype=torch.bool))

    materialized, materialized_prior = run(
        torch.full((num_rows,), 2).cuda(), torch.device("cuda")
    )
    assert torch.equal(materialized, expected)
    assert torch.equal(materialized_prior, expected_prior)

    # Neighbouring steps are the value a row reads when the broadcast stride is
    # ignored; the production tensor reads uninitialized memory instead.
    neighbours = torch.tensor([2, 0, 0, 0]).cuda()
    broadcast, broadcast_prior = run(
        neighbours[0].expand(num_rows), torch.device("cuda")
    )
    assert torch.equal(broadcast, expected)
    assert torch.equal(broadcast_prior, expected_prior)

    step = torch.tensor(2).cuda()
    production, production_prior = run(step.expand(num_rows), torch.device("cuda"))
    assert torch.equal(production, expected)
    assert torch.equal(production_prior, expected_prior)


def test_draft_sampler_uses_draft_key_and_advances_context(monkeypatch):
    class StubSpeculator(DraftModelSpeculator):
        def capture(self): ...

        def init_cudagraph_manager(self, cudagraph_mode): ...

        def load_draft_model(self, target_model, target_attn_layer_names): ...

        def propose(self, *args, **kwargs): ...

    class StubModel:
        @staticmethod
        def compute_logits(hidden_states):
            return torch.zeros(hidden_states.shape[0], 8)

    class StubWatermarker:
        context_width = 2

        @staticmethod
        def sample(logits, contexts, random_sampler=None, skip_mask=None):
            assert random_sampler is not None
            return WatermarkSample(
                torch.where(skip_mask, random_sampler(logits), 7), logits
            )

    speculator = object.__new__(StubSpeculator)
    speculator.model = StubModel()
    speculator.use_fp64_gumbel = False
    speculator.acceptance_estimator = None
    draft_watermarker = DraftWatermarker(
        StubWatermarker(),
        max_num_reqs=2,
        device=torch.device("cpu"),
        num_speculative_steps=1,
        deduplicate_contexts="none",
        deduplicate_contexts_max_history=None,
    )
    draft_watermarker.contexts.copy_(torch.tensor([[1, 2], [3, 4]]))
    draft_watermarker.enabled = torch.tensor([True, False])
    speculator.draft_watermarker = draft_watermarker
    monkeypatch.setattr(
        "vllm.v1.watermarking.watermarker.gumbel_sample",
        lambda *args, **kwargs: torch.tensor([3, 4]),
    )

    sampled = speculator.sample_draft(
        hidden_states=torch.zeros(2, 4),
        sample_src_positions=torch.zeros(2, dtype=torch.int64),
        idx_mapping=torch.tensor([0, 1]),
        temperature=torch.ones(2),
        seeds=torch.zeros(2, dtype=torch.int64),
        draft_step=torch.tensor(0),
        draft_logits=torch.zeros(2, 1, 8),
    )

    assert torch.equal(sampled, torch.tensor([7, 4]))
    assert torch.equal(draft_watermarker.contexts, torch.tensor([[2, 7], [4, 4]]))


def test_draft_sampler_uses_ordinary_samples_for_repeated_contexts(monkeypatch):
    class StubWatermarker:
        context_width = 1

        def __init__(self):
            self.samples = iter((2, 1, 8))

        def sample(self, logits, contexts, random_sampler=None, skip_mask=None):
            watermarked = torch.tensor([next(self.samples)])
            return WatermarkSample(
                torch.where(skip_mask, random_sampler(logits), watermarked), logits
            )

    draft_watermarker = DraftWatermarker(
        StubWatermarker(),
        max_num_reqs=1,
        device=torch.device("cpu"),
        num_speculative_steps=3,
        deduplicate_contexts="single_turn",
        deduplicate_contexts_max_history=None,
    )
    draft_watermarker.prepare(
        contexts=torch.tensor([[1]]),
        enabled=torch.tensor([True]),
        all_token_ids=torch.tensor([[100, 1, 0, 0]]),
        prompt_lens=torch.tensor([1]),
        total_lens=torch.tensor([2]),
    )
    idx_mapping = torch.tensor([0])
    temperature = torch.tensor([1.0])
    ordinary_samples = iter((torch.tensor([5]), torch.tensor([6]), torch.tensor([4])))
    monkeypatch.setattr(
        "vllm.v1.watermarking.watermarker.gumbel_sample",
        lambda *args, **kwargs: next(ordinary_samples),
    )
    seeds = torch.tensor([0])
    draft_logits = torch.zeros(1, 3, 8)

    first = draft_watermarker.sample(
        torch.zeros(1, 8),
        idx_mapping,
        temperature=temperature,
        seed=seeds,
        pos=torch.tensor([0]),
        apply_temperature=True,
        is_drafting=True,
        logits_cache_col=torch.tensor(0),
        logits_cache=draft_logits,
        use_fp64=False,
    )
    second = draft_watermarker.sample(
        torch.zeros(1, 8),
        idx_mapping,
        temperature=temperature,
        seed=seeds,
        pos=torch.tensor([1]),
        apply_temperature=True,
        is_drafting=True,
        logits_cache_col=torch.tensor(1),
        logits_cache=draft_logits,
        use_fp64=False,
    )
    repeated = draft_watermarker.sample(
        torch.zeros(1, 8),
        idx_mapping,
        temperature=temperature,
        seed=seeds,
        pos=torch.tensor([2]),
        apply_temperature=True,
        is_drafting=True,
        logits_cache_col=torch.tensor(2),
        logits_cache=draft_logits,
        use_fp64=False,
    )

    assert first.item() == 2
    assert second.item() == 1
    assert repeated.item() == 4


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
def test_draft_context_deduplication_reads_strided_request_indices():
    """DSpark samples a step at a time, so req_indices is an `idx_map[:, i]` column."""
    num_reqs = 4
    num_speculative_steps = 3
    all_token_ids = torch.tensor(
        [[3, 4, 3, 4], [9, 9, 9, 9], [3, 4, 3, 4], [9, 9, 9, 9]]
    ).cuda()
    prompt_lens = torch.zeros(num_reqs, dtype=torch.int64).cuda()
    total_lens = torch.full((num_reqs,), 4).cuda()
    contexts = torch.tensor([[3, 4]]).expand(num_reqs, 2).contiguous().cuda()
    steps = torch.zeros(num_reqs, dtype=torch.int64).cuda()
    enabled = torch.ones(num_reqs, dtype=torch.bool).cuda()

    idx_map = (
        torch.arange(num_reqs)
        .repeat_interleave(num_speculative_steps)
        .view(num_reqs, num_speculative_steps)
        .cuda()
    )
    column = idx_map[:, 1]
    assert column.stride(0) == num_speculative_steps

    def run(req_indices):
        prior = torch.zeros(
            num_reqs, num_speculative_steps, 2, dtype=torch.int64, device="cuda"
        )
        return draft_watermarking_mask(
            all_token_ids,
            req_indices,
            prompt_lens,
            total_lens,
            prior,
            contexts,
            steps,
            enabled,
            max_history=None,
            include_prompt=False,
        ).cpu()

    expected = torch.tensor([False, True, False, True])
    assert torch.equal(run(column.contiguous()), expected)
    assert torch.equal(run(column), expected)


def test_draft_sampler_deduplicates_against_committed_history(monkeypatch):
    class StubWatermarker:
        context_width = 1

        @staticmethod
        def sample(logits, contexts, random_sampler=None, skip_mask=None):
            return WatermarkSample(
                torch.where(skip_mask, random_sampler(logits), 7), logits
            )

    draft_watermarker = DraftWatermarker(
        StubWatermarker(),
        max_num_reqs=1,
        device=torch.device("cpu"),
        num_speculative_steps=1,
        deduplicate_contexts="single_turn",
        deduplicate_contexts_max_history=None,
    )
    draft_watermarker.prepare(
        contexts=torch.tensor([[1]]),
        enabled=torch.tensor([True]),
        all_token_ids=torch.tensor([[100, 1, 2, 1]]),
        prompt_lens=torch.tensor([1]),
        total_lens=torch.tensor([4]),
    )

    monkeypatch.setattr(
        "vllm.v1.watermarking.watermarker.gumbel_sample",
        lambda *args, **kwargs: torch.tensor([4]),
    )
    sampled = draft_watermarker.sample(
        torch.zeros(1, 8),
        torch.tensor([0]),
        temperature=torch.tensor([1.0]),
        seed=torch.tensor([0]),
        pos=torch.tensor([0]),
        apply_temperature=True,
        is_drafting=True,
        logits_cache_col=torch.tensor(0),
        logits_cache=torch.zeros(1, 1, 8),
        use_fp64=False,
    )

    assert sampled.item() == 4


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
@pytest.mark.parametrize("idx_mapping_layout", ["contiguous", "column_view"])
def test_draft_watermarker_deduplicates_across_a_block_on_cuda(
    monkeypatch, idx_mapping_layout
):
    """Drive the whole draft sampler over a block with the production shapes.

    `draft_step` is a 0-d tensor, and a speculator that samples one step at a
    time hands over a column of its index map. Each step is checked against the
    committed history the window still reaches and against the earlier steps of
    the block.
    """
    num_reqs, num_steps, vocab_size = 2, 3, 16
    watermarker = GumbelWatermarker(key=42, context_width=2)
    draft_watermarker = DraftWatermarker(
        watermarker,
        max_num_reqs=num_reqs,
        device=torch.device("cuda"),
        num_speculative_steps=num_steps,
        deduplicate_contexts="single_turn",
        deduplicate_contexts_max_history=4,
    )
    draft_watermarker.prepare(
        contexts=torch.tensor([[5, 6], [6, 3]]).cuda(),
        enabled=torch.ones(num_reqs, dtype=torch.bool).cuda(),
        all_token_ids=torch.tensor([[5, 6, 7, 5, 6], [8, 9, 6, 6, 3]]).cuda(),
        prompt_lens=torch.zeros(num_reqs, dtype=torch.int64).cuda(),
        total_lens=torch.full((num_reqs,), 5).cuda(),
    )

    if idx_mapping_layout == "contiguous":
        idx_mapping = torch.arange(num_reqs, dtype=torch.int32).cuda()
    else:
        idx_mapping = (
            torch.arange(num_reqs, dtype=torch.int32)
            .repeat_interleave(num_steps)
            .view(num_reqs, num_steps)
            .cuda()
        )[:, 1]
        assert idx_mapping.stride(0) == num_steps

    skip_masks: list[torch.Tensor] = []
    original = philox_gumbel_sample

    def spy(*args, **kwargs):
        skip_masks.append(kwargs["skip_mask"].cpu())
        return original(*args, **kwargs)

    monkeypatch.setattr("vllm.v1.worker.gpu.sample.watermark.philox_gumbel_sample", spy)

    draft_logits = torch.zeros(num_reqs, num_steps, vocab_size + 1).cuda()
    step_logits = []
    for step in range(num_steps):
        # Token 6 wins by a margin no Gumbel draw can close, on either path.
        logits = torch.zeros(num_reqs, vocab_size).cuda()
        logits[:, 6] = 100.0
        logits[:, 0] = float(step) + 1.0
        step_logits.append(logits)
        sampled = draft_watermarker.sample(
            logits,
            idx_mapping,
            temperature=torch.ones(num_reqs).cuda(),
            seed=torch.zeros(num_reqs, dtype=torch.int64).cuda(),
            pos=torch.full((num_reqs,), 5 + step, dtype=torch.int64).cuda(),
            apply_temperature=True,
            is_drafting=True,
            logits_cache_col=torch.tensor(step).cuda(),
            logits_cache=draft_logits,
            use_fp64=False,
        )
        assert torch.equal(sampled.cpu(), torch.tensor([6, 6]))

    # Step 0: request 0's [5, 6] is committed at position 2, request 1's [6, 3]
    # is new. Step 1: [6, 6] and [3, 6] are both new. Step 2: request 0's
    # [6, 6] repeats its own step-1 context and request 1's repeats position 4.
    assert [mask.tolist() for mask in skip_masks] == [
        [True, False],
        [False, False],
        [True, True],
    ]
    assert torch.equal(
        draft_watermarker.prior_contexts[:num_reqs].cpu(),
        torch.tensor([[[5, 6], [6, 6], [6, 6]], [[6, 3], [3, 6], [6, 6]]]),
    )
    assert torch.equal(
        draft_watermarker.contexts[:num_reqs].cpu(), torch.tensor([[6, 6], [6, 6]])
    )
    for step, logits in enumerate(step_logits):
        # The 0-d draft step selects the logits cache column.
        assert torch.equal(draft_logits[:, step, :vocab_size], logits)
    assert not draft_logits[:, :, -1].any()


@pytest.mark.skipif(
    not current_platform.is_cuda_alike(), reason="requires a CUDA-like accelerator"
)
def test_draft_watermarker_cuda_graph_reads_current_prompt_lens():
    draft_watermarker = DraftWatermarker(
        GumbelWatermarker(key=42, context_width=2),
        max_num_reqs=1,
        device=torch.device("cuda"),
        num_speculative_steps=1,
        deduplicate_contexts="single_turn",
        deduplicate_contexts_max_history=None,
    )
    prompt_lens = UvaBackedTensor(1, dtype=torch.int32)
    contexts = torch.tensor([[5, 6]]).cuda()
    enabled = torch.ones(1, dtype=torch.bool).cuda()
    # [5, 6] repeats only while the prompt counts as history.
    all_token_ids = torch.tensor([[5, 6, 9, 5, 6]]).cuda()
    total_lens = torch.tensor([5]).cuda()
    logits = torch.zeros(1, 8).cuda()
    idx_mapping = torch.zeros(1, dtype=torch.int32).cuda()
    temperature = torch.ones(1).cuda()
    draft_step = torch.tensor(0).cuda()

    def prepare():
        draft_watermarker.prepare(
            contexts, enabled, all_token_ids, prompt_lens.gpu, total_lens
        )

    def watermark_mask():
        _, mask = draft_watermarker._sampling_state(
            logits, idx_mapping, temperature, draft_step, contexts
        )
        return mask

    prepare()
    assert not watermark_mask().item()
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        mask = watermark_mask()
    # A step that adds requests moves prompt_lens.gpu to the next pool buffer.
    prompt_lens.np[:] = 3
    prompt_lens.copy_to_uva()
    prepare()
    graph.replay()
    assert mask.item()


def test_dspark_reduced_vocab_draft_sampler_applies_watermarking(monkeypatch):
    speculator = object.__new__(DSparkSpeculator)
    speculator.draft_logits = torch.zeros(2, 1, 8)
    speculator._d2t_scatter_index = torch.tensor([1, 5])
    speculator._draft_scatter_buf = torch.full((2, 8), float("-inf"))
    speculator.temperature = torch.ones(2)
    speculator.seeds = torch.zeros(2, dtype=torch.int64)
    speculator._step_cols = torch.tensor([0])
    speculator.use_fp64_gumbel = False
    speculator.acceptance_estimator = None
    watermark_logits: list[torch.Tensor] = []

    def sample(
        logits,
        idx_mapping,
        temperature,
        seed,
        pos,
        apply_temperature,
        is_drafting,
        logits_cache=None,
        logits_cache_col=None,
        use_fp64=False,
        logits_cache_source=None,
    ):
        watermark_logits.append(logits.clone())
        return torch.tensor([4, 5])

    speculator.draft_watermarker = SimpleNamespace(sample=sample)

    sampled = speculator._sample_logits(
        torch.tensor([[10.0, 20.0], [30.0, 40.0]]),
        torch.tensor([0, 1]),
        torch.tensor([1, 1]),
        0,
    )

    assert torch.equal(sampled, torch.tensor([4, 5]))
    assert torch.equal(
        watermark_logits[0][:, [1, 5]], torch.tensor([[10, 20], [30, 40]])
    )
    assert torch.isneginf(watermark_logits[0][:, [0, 2, 3, 4, 6, 7]]).all()


def test_dspark_target_only_watermarking_leaves_drafts_unwatermarked(monkeypatch):
    speculator = object.__new__(DSparkSpeculator)
    speculator.draft_logits = torch.zeros(2, 1, 8)
    speculator._d2t_scatter_index = None
    speculator.temperature = torch.ones(2)
    speculator.seeds = torch.zeros(2, dtype=torch.int64)
    speculator._step_cols = torch.tensor([0])
    speculator.use_fp64_gumbel = False
    speculator.acceptance_estimator = None
    speculator.draft_watermarker = None
    monkeypatch.setattr(
        "vllm.v1.worker.gpu.spec_decode.dspark.speculator.gumbel_sample",
        lambda *args, **kwargs: torch.tensor([3, 4]),
    )

    sampled = speculator._sample_logits(
        torch.zeros(2, 8),
        torch.tensor([0, 1]),
        torch.tensor([1, 1]),
        0,
    )

    assert torch.equal(sampled, torch.tensor([3, 4]))
