From 8a5cf1ccd65e8ac7635c402c1ec0b08988bc26ca Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Juan=20P=C3=A9rez=20de=20Algaba?= <124347725+jperezdealgaba@users.noreply.github.com> Date: Wed, 10 Jun 2026 11:31:43 +0200 Subject: [PATCH] [Security] Fix remote DoS via invalid recovered token reinjection (#44744) Signed-off-by: jperezde --- tests/v1/sample/test_rejection_sampler.py | 93 +++++++++++++++++++++++ vllm/v1/sample/rejection_sampler.py | 7 +- 2 files changed, 99 insertions(+), 1 deletion(-) diff --git a/tests/v1/sample/test_rejection_sampler.py b/tests/v1/sample/test_rejection_sampler.py index e628f903792..10c4d448f7f 100644 --- a/tests/v1/sample/test_rejection_sampler.py +++ b/tests/v1/sample/test_rejection_sampler.py @@ -998,6 +998,99 @@ def test_sample_recovered_tokens_uses_fp64_exponential_race_when_requested(): assert torch.equal(actual, expected) +@pytest.mark.parametrize("no_draft_probs", [True, False]) +@pytest.mark.parametrize( + "vocab_size", + [ + 100, # below BLOCK_SIZE: single partial tile with many padding entries + 8193, # BLOCK_SIZE + 1: only 1 valid entry in the last tile + 10000, # non-aligned, moderate tail + 151936, # real-world Qwen3 vocab size from the CVE report + ], +) +def test_sample_recovered_tokens_vocab_boundary(vocab_size: int, no_draft_probs: bool): + """Regression test for GHSA-8wr5-jm2h-8r4f. + + When vocab_size is not a multiple of BLOCK_SIZE (8192), the last Triton + tile extends beyond the vocabulary. If all valid entries in that tail tile + have zero target probability, the out-of-range masked positions (score 0) + could win the tl.max tie-break, producing recovered_id >= vocab_size. + This test forces that scenario and asserts every recovered token is valid. + """ + BLOCK_SIZE = 8192 + batch_size = 2 + max_spec_len = 3 + num_tokens = batch_size * max_spec_len + + last_tile_start = (vocab_size // BLOCK_SIZE) * BLOCK_SIZE + + target_probs = torch.rand( + num_tokens, vocab_size, dtype=torch.float32, device=DEVICE_TYPE + ) + if last_tile_start > 0: + # Zero out valid entries in the last partial tile so the only + # non-zero scores come from earlier, fully-covered tiles. + target_probs[:, last_tile_start:] = 0.0 + else: + # vocab_size < BLOCK_SIZE: single tile. Concentrate all mass on + # entry 0 so the NO_DRAFT_PROBS path (which zeroes the draft + # token entry) can drive all valid scores to zero. + target_probs = torch.zeros_like(target_probs) + target_probs[:, 0] = 1.0 + # Re-normalize so it's a valid distribution. + target_probs = target_probs / target_probs.sum(dim=-1, keepdim=True) + + draft_probs = torch.rand( + num_tokens, vocab_size, dtype=torch.float32, device=DEVICE_TYPE + ) + draft_probs = torch.nn.functional.softmax(draft_probs, dim=-1) + + if last_tile_start == 0: + # Force draft token to 0 so the NO_DRAFT_PROBS path zeroes the + # only non-zero entry, leaving all valid scores at zero. + draft_token_ids = torch.zeros( + num_tokens, 1, dtype=torch.int32, device=DEVICE_TYPE + ) + else: + draft_token_ids = torch.randint( + 0, vocab_size, (num_tokens, 1), dtype=torch.int32, device=DEVICE_TYPE + ) + + temperature = torch.ones(batch_size, dtype=torch.float32, device=DEVICE_TYPE) + generators = { + i: torch.Generator(device=DEVICE_TYPE).manual_seed(42 + i) + for i in range(batch_size) + } + sampling_metadata = create_sampling_metadata( + all_greedy=False, temperature=temperature, generators=generators + ) + + spec_decode_metadata = create_spec_decode_metadata( + draft_token_ids.reshape(batch_size, max_spec_len).tolist(), + torch.rand(num_tokens, vocab_size, device=DEVICE_TYPE), + ) + + recovered = sample_recovered_tokens( + max_spec_len, + spec_decode_metadata.num_draft_tokens, + spec_decode_metadata.cu_num_draft_tokens, + draft_token_ids.squeeze(-1), + None if no_draft_probs else draft_probs, + target_probs, + sampling_metadata, + device=DEVICE_TYPE, + ) + + assert (recovered >= 0).all(), ( + f"Recovered token IDs contain negative values: " + f"{recovered[recovered < 0].tolist()}" + ) + assert (recovered < vocab_size).all(), ( + f"Recovered token IDs >= vocab_size ({vocab_size}): " + f"{recovered[recovered >= vocab_size].tolist()}" + ) + + ########################### Tests for Synthetic Rejection Sampling ######### diff --git a/vllm/v1/sample/rejection_sampler.py b/vllm/v1/sample/rejection_sampler.py index 153677e35fa..1c1e57427f3 100644 --- a/vllm/v1/sample/rejection_sampler.py +++ b/vllm/v1/sample/rejection_sampler.py @@ -921,12 +921,17 @@ def sample_recovered_tokens_kernel( other=0.0, ) - # Local tile reduction + # Local tile reduction. + # Mask out-of-vocabulary entries to -inf so they can never win + # the argmax — prevents producing recovered_id >= vocab_size + # when all valid entries in the last tile have zero probability. score = prob * inv_q + score = tl.where(vocab_mask, score, float("-inf")) local_max, local_id = tl.max(score, axis=0, return_indices=True) if local_max > max_val: max_val = local_max recovered_id = v + local_id + recovered_id = tl.minimum(recovered_id, vocab_size - 1) tl.store(output_token_ids_ptr + token_idx, recovered_id)