Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
294 changes: 288 additions & 6 deletions tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -835,9 +835,27 @@ def test_fence_at_update_state_after_alloc(request_runner):
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.run(decoded_tokens=[EOS_TOKEN_ID], complete_transfers=False)

# Capture fence snapshots to verify block 0 is registered.
fence_snapshots: list[dict] = []

def capture_fence():
fence_snapshots.append(
dict(runner.connector_scheduler._block_id_to_pending_jobs)
)

runner.run(
decoded_tokens=[EOS_TOKEN_ID],
complete_transfers=False,
post_step_fn=capture_fence,
)
assert runner.connector_scheduler._block_id_to_pending_jobs

# Verify fence was populated with the store job's block IDs.
populated_fence = next((f for f in fence_snapshots if f), None)
assert populated_fence is not None, "Fence was never populated"
assert len(populated_fence) > 0, "Fence is empty"

runner.scheduler.reset_prefix_cache()
runner.new_request(token_ids=[0] * 4)
runner.connector_scheduler._maximal_prefix_lookup = lambda key, req_context: 1
Expand Down Expand Up @@ -868,9 +886,27 @@ def test_fence_at_build_store_jobs(request_runner):
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.run(decoded_tokens=[EOS_TOKEN_ID], complete_transfers=False)

# Capture fence snapshots to verify block 0 is registered.
fence_snapshots: list[dict] = []

def capture_fence():
fence_snapshots.append(
dict(runner.connector_scheduler._block_id_to_pending_jobs)
)

runner.run(
decoded_tokens=[EOS_TOKEN_ID],
complete_transfers=False,
post_step_fn=capture_fence,
)
assert runner.connector_scheduler._block_id_to_pending_jobs

# Verify fence was populated with the store job's block IDs.
populated_fence = next((f for f in fence_snapshots if f), None)
assert populated_fence is not None, "Fence was never populated"
assert len(populated_fence) > 0, "Fence is empty"

runner.scheduler.reset_prefix_cache()
runner.new_request(token_ids=[1] * 4)
runner.connector_scheduler._maximal_prefix_lookup = lambda key, req_context: 0
Expand Down Expand Up @@ -950,8 +986,8 @@ def setup(r, max_offload_tokens):
token_ids=[0] * offloaded_block_size * 3,
kv_transfer_params={"max_offload_tokens": max_offload_tokens},
)
r.manager.prepare_store.side_effect = (
lambda keys, req_context: generate_store_output(keys)
r.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)

# Pending offloads drain via non-blocking stepping, not a flush, so no
Expand Down Expand Up @@ -1049,8 +1085,8 @@ def test_offload_prompt_only(request_runner, async_scheduling: bool):
extra_config_overrides={"offload_prompt_only": True},
)

runner.manager.prepare_store.side_effect = (
lambda keys, req_context: generate_store_output(keys)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)

runner.new_request(token_ids=[0] * offloaded_block_size * num_prompt_blocks)
Expand Down Expand Up @@ -2053,3 +2089,249 @@ def test_full_attn_store_then_load(self, request_runner, async_scheduling: bool)
(1, 1),
),
)


# ---------------------------------------------------------------------------
# Tests for request_finished fence population with in-flight pending stores.
# ---------------------------------------------------------------------------


def test_request_finished_with_pending_stores_populates_fence(request_runner):
"""When a request finishes with in-flight store jobs, the fence index
(_block_id_to_pending_jobs) is correctly populated with the store jobs'
non_sliding_window_block_ids.

This prevents data corruption when a subsequent request reuses the same
GPU blocks before the store completes.
"""
block_size = 4
block_size_factor = 1
offloaded_block_size = block_size * block_size_factor

# Use 2 GPU blocks so the second run reuses the same blocks,
# triggering a fence-based flush of the in-flight job from run 1.
runner = request_runner(
block_size=block_size,
num_gpu_blocks=2,
async_scheduling=False,
block_size_factor=block_size_factor,
)

# 4 prompt tokens → 1 GPU block (block 0)
runner.new_request(token_ids=[0] * offloaded_block_size)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)

# Capture fence state at each step to verify it was populated.
fence_snapshots: list[dict] = []
job_block_ids: set[int] = set()

def capture_fence():
fence_snapshots.append(
dict(runner.connector_scheduler._block_id_to_pending_jobs)
)
for js in runner.connector_scheduler._jobs.values():
if js.is_store:
job_block_ids.update(js.non_sliding_window_block_ids or [])

# Run 1: create store job, finish request, populate fence.
# With non-blocking drain (#45595), the job stays in-flight.
runner.run(
decoded_tokens=[EOS_TOKEN_ID],
complete_transfers=False,
post_step_fn=capture_fence,
)

# Verify fence was populated at some point during the run.
assert len(job_block_ids) > 0, "No store job was created"
populated_fence = next((f for f in fence_snapshots if len(f) > 0), None)
assert populated_fence is not None, "Fence was never populated"

# Verify fence contained the job's non-SW block IDs.
for bid in job_block_ids:
assert bid in populated_fence, f"Block {bid} not in fence: {populated_fence}"

# Run 2: block reuse triggers fence-based flush → cleanup.
runner.scheduler.reset_prefix_cache()
runner.new_request(token_ids=[0] * offloaded_block_size)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.run(
decoded_tokens=[EOS_TOKEN_ID],
expected_stored=(0,),
expected_flushed=(0,),
)

# Verify fence is empty after full lifecycle (cleanup happened).
assert runner.connector_scheduler._block_id_to_pending_jobs == {}
# req_status should be removed.
req_id = str(runner.req_id)
assert req_id not in runner.connector_scheduler._req_status


def test_multiple_in_flight_stores_all_flushed_by_fence(request_runner):
"""When a request finishes with multiple in-flight store jobs,
ALL jobs are flushed when a new request reuses their blocks.

Uses three runner.run() calls:
- Run 1: decode fills a block → job_0 created
- Run 2: decode fills another block + EOS → job_1 created, request finishes
- Run 3: block reuse → both jobs flushed via fence
"""
block_size = 4
block_size_factor = 1
offloaded_block_size = block_size * block_size_factor

# 4 GPU blocks: block 0 is null, blocks 1-3 are usable.
runner = request_runner(
block_size=block_size,
num_gpu_blocks=4,
async_scheduling=False,
block_size_factor=block_size_factor,
)

# Prompt: 4 tokens → block 1
runner.new_request(token_ids=[0] * offloaded_block_size)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)

# Run 1: 4 decoded tokens → block 2 full → job_0 created for block 1.
runner.run(
decoded_tokens=[0] * offloaded_block_size,
complete_transfers=False,
)
assert len(runner.connector_scheduler._jobs) >= 1

# Run 2: 4 more tokens + EOS → block 3 full → more jobs created.
# Request finishes → all jobs registered in fence.
runner.run(
decoded_tokens=[0] * offloaded_block_size + [EOS_TOKEN_ID],
complete_transfers=False,
)
num_jobs = len(runner.connector_scheduler._jobs)
assert num_jobs >= 2, f"Expected multiple in-flight jobs, got {num_jobs}"

# Run 3: block reuse → fence flushes both jobs.
runner.scheduler.reset_prefix_cache()
runner.new_request(token_ids=[0] * offloaded_block_size * 3)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.run(
decoded_tokens=[EOS_TOKEN_ID],
expected_stored=(0, 1, 2),
expected_flushed=(0, 1, 2),
)

# Post-condition: fence cleaned up, all jobs gone.
assert runner.connector_scheduler._block_id_to_pending_jobs == {}
assert len(runner.connector_scheduler._jobs) == 0


def test_request_finished_mixed_full_attn_and_sliding_window(
request_runner,
):
"""With both FullAttention and SlidingWindow groups, a single store job
has both non_sliding_window_block_ids and sliding_window_block_ids.

request_finished only registers non-SW blocks in the fence.
SW blocks were already registered at store creation time.
"""
block_size = 4
sliding_window = 8 # 2 blocks

kv_cache_groups = [
KVCacheGroupSpec(
["layer0"],
FullAttentionSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["layer1"],
SlidingWindowSpec(
block_size=block_size,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
sliding_window=sliding_window,
),
),
]

# Use 4 GPU blocks (2 per group) so run 2 reuses the same blocks,
# triggering a fence-based flush.
runner = request_runner(
block_size=block_size,
num_gpu_blocks=4,
async_scheduling=False,
kv_cache_groups=kv_cache_groups,
)

# 1 block of prompt (4 tokens) — 1 block per group.
runner.new_request(token_ids=[0] * block_size)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)

# Capture fence state and job block IDs at each step.
fence_snapshots: list[dict] = []
sw_block_ids: set[int] = set()
non_sw_block_ids: set[int] = set()

def capture_fence():
fence_snapshots.append(
dict(runner.connector_scheduler._block_id_to_pending_jobs)
)
for js in runner.connector_scheduler._jobs.values():
if js.is_store:
sw_block_ids.update(js.sliding_window_block_ids or [])
non_sw_block_ids.update(js.non_sliding_window_block_ids or [])

# Run 1: create store job, finish request, populate fence.
runner.run(
decoded_tokens=[EOS_TOKEN_ID],
complete_transfers=False,
post_step_fn=capture_fence,
)

# Verify job had both SW and non-SW blocks.
assert len(sw_block_ids) > 0, "No SW blocks in store job"
assert len(non_sw_block_ids) > 0, "No non-SW blocks in store job"

# Find the fence snapshot where both SW and non-SW blocks were present.
# SW blocks should appear at creation time, non-SW at request_finished.
populated_fence = None
for fence in fence_snapshots:
has_sw = all(bid in fence for bid in sw_block_ids)
has_non_sw = all(bid in fence for bid in non_sw_block_ids)
if has_sw and has_non_sw:
populated_fence = fence
break

assert populated_fence is not None, (
f"Fence never contained both SW {sw_block_ids} and "
f"non-SW {non_sw_block_ids} blocks. Snapshots: {fence_snapshots}"
)

# Run 2: block reuse triggers fence-based flush of the old job.
runner.scheduler.reset_prefix_cache()
runner.new_request(token_ids=[0] * block_size)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.run(
decoded_tokens=[EOS_TOKEN_ID],
expected_stored=((0, 0), (1, 0)),
expected_flushed=((1, 0),),
)

# Verify fence is empty after full lifecycle (cleanup happened).
assert runner.connector_scheduler._block_id_to_pending_jobs == {}
assert len(runner.connector_scheduler._jobs) == 0
17 changes: 14 additions & 3 deletions tests/v1/kv_connector/unit/offloading_connector/utils.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Iterable, Iterator
from collections.abc import Callable, Iterable, Iterator
from dataclasses import dataclass
from typing import Any
from unittest.mock import MagicMock
Expand Down Expand Up @@ -430,14 +430,21 @@ def _update_gpu_blocks(self):
for block_idx, block in enumerate(blocks):
self.gpu_blocks[block.block_id] = GPUBlock(group_idx, block_idx)

def _run(self, decoded_tokens: list[int], complete_transfers: bool):
def _run(
self,
decoded_tokens: list[int],
complete_transfers: bool,
post_step_fn: Callable[[], None] | None = None,
):
"""
Runs multiple engine (scheduler + worker) steps.
Assumes a single request is running.

Args:
decoded_tokens: the tokens to yield at each step.
complete_transfers: complete transfers immediately
post_step_fn: optional callback invoked after each step's
update_from_output(), before the next schedule().
"""

tokens_iter = iter(decoded_tokens)
Expand Down Expand Up @@ -500,6 +507,9 @@ def _run(self, decoded_tokens: list[int], complete_transfers: bool):
else:
self.scheduler.update_from_output(scheduler_output, model_runner_output)

if post_step_fn is not None:
post_step_fn()

if (
prev_token_id == EOS_TOKEN_ID
and prev_token_id != token_id
Expand Down Expand Up @@ -545,6 +555,7 @@ def run(
expected_stored: tuple[int | tuple[int, int], ...] = (),
expected_loaded: tuple[int | tuple[int, int], ...] = (),
expected_flushed: tuple[int | tuple[int, int], ...] = (),
post_step_fn: Callable[[], None] | None = None,
):
"""
Runs multiple engine (scheduler + worker) steps.
Expand All @@ -570,7 +581,7 @@ def run(
expected_flushed_gpu_blocks = self._to_gpu_blocks(expected_flushed)

self.manager.reset_mock()
self._run(decoded_tokens, complete_transfers)
self._run(decoded_tokens, complete_transfers, post_step_fn=post_step_fn)

loaded_gpu_blocks: set[GPUBlock] = set()
for transfer in self.completed_loads:
Expand Down
Loading