Skip to content
Draft
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
8 changes: 8 additions & 0 deletions tests/v1/kv_connector/unit/test_mooncake_store_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from vllm.distributed.kv_events import BlockStored
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorRole,
KVConnectorSchedulerContext,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store import (
connector as mooncake_store_connector,
Expand Down Expand Up @@ -88,6 +89,13 @@ def test_scheduler_role_initializes_store_scheduler_only():
mock_worker.assert_not_called()
assert connector.connector_scheduler is mock_scheduler.return_value
assert connector.connector_worker is None
block_pool = MagicMock()

scheduler_context = KVConnectorSchedulerContext(block_pool, MagicMock())
connector.bind_scheduler_context(scheduler_context)
mock_scheduler.return_value.bind_scheduler_context.assert_called_once_with(
scheduler_context
)


def test_worker_methods_delegate_to_store_worker():
Expand Down
154 changes: 102 additions & 52 deletions tests/v1/kv_connector/unit/test_mooncake_store_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,13 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from types import SimpleNamespace
from unittest.mock import MagicMock

import pytest

from vllm.distributed.kv_transfer.kv_connector.v1.base import (
KVConnectorSchedulerContext,
)
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.data import (
LoadSpec,
MooncakeStoreWorkerMetadata,
Expand All @@ -14,7 +18,9 @@
from vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store.scheduler import (
MooncakeStoreScheduler,
)
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
from vllm.v1.core.block_pool import BlockPool
from vllm.v1.core.kv_cache_manager import KVCacheBlocks


def _make_bare_scheduler(
Expand All @@ -37,15 +43,47 @@ def _make_bare_scheduler(
scheduler._unfinished_request_ids = {"req-0"}
scheduler._unfinished_requests = {}
scheduler._request_trackers = {}
scheduler._gpu_block_pool = BlockPool(
gpu_block_pool = BlockPool(
num_gpu_blocks=64, enable_caching=True, hash_block_size=hash_block_size
)
scheduler._num_workers = 1
scheduler._next_store_job_id = 0
scheduler._pinned_saves = {}
scheduler._test_current_block_ids = {}
scheduler._test_scheduler_context_calls = []

def get_request_blocks(request_id: str) -> KVCacheBlocks:
scheduler._test_scheduler_context_calls.append(request_id)
block_ids = scheduler._test_current_block_ids[request_id]
return KVCacheBlocks(
tuple(
[gpu_block_pool.blocks[block_id] for block_id in group]
for group in block_ids
)
)

scheduler.bind_scheduler_context(
KVConnectorSchedulerContext(gpu_block_pool, get_request_blocks)
)
assert scheduler._scheduler_context is not None
assert scheduler._scheduler_context.gpu_block_pool is gpu_block_pool
return scheduler


def _get_gpu_block_pool(scheduler: MooncakeStoreScheduler) -> BlockPool:
assert scheduler._scheduler_context is not None
return scheduler._scheduler_context.gpu_block_pool


def _set_unfinished_request(
scheduler: MooncakeStoreScheduler,
request: SimpleNamespace,
block_ids: tuple[list[int], ...],
) -> None:
scheduler._unfinished_requests["req-0"] = request
scheduler._test_current_block_ids["req-0"] = block_ids


def _make_scheduler_output(*, scheduled_spec_tokens: list[int] | None):
return SimpleNamespace(
finished_req_ids=set(),
Expand Down Expand Up @@ -108,11 +146,23 @@ def _make_new_scheduler_output() -> SimpleNamespace:
)


def test_update_state_after_alloc_does_not_materialize_block_ids():
scheduler = _make_bare_scheduler()
request = SimpleNamespace(request_id="req-0")
blocks = MagicMock(spec=KVCacheBlocks)

scheduler.update_state_after_alloc(request, blocks, num_external_tokens=0)

blocks.get_block_ids.assert_not_called()
assert scheduler._unfinished_requests["req-0"] is request


def test_scheduler_only_tracks_token_ids_for_kv_events():
for enable_kv_events in (False, True):
scheduler = _make_bare_scheduler()
scheduler.enable_kv_events = enable_kv_events
scheduler._unfinished_requests["req-0"] = (
_set_unfinished_request(
scheduler,
_make_new_scheduler_output().request,
([0, 1],),
)
Expand Down Expand Up @@ -168,11 +218,10 @@ def _add_unfinished_request(
block_hashes=block_hashes,
num_output_placeholders=0,
)
scheduler._unfinished_requests["req-0"] = (request, ([0, 1],))
_set_unfinished_request(scheduler, request, ([0, 1, 2],))
scheduler._request_trackers["req-0"] = RequestTracker(
req_id="req-0",
token_len=44,
allocated_block_ids=([0, 1],),
num_saved_tokens=32,
token_ids=token_ids[:44],
prefill_end_tokens=prefill_end_tokens,
Expand All @@ -197,7 +246,6 @@ def _setup_decode_request(
)
tracker = scheduler._request_trackers["req-0"]
tracker.token_len = token_len
tracker.allocated_block_ids = ([0, 1, 2],)
tracker.token_ids = token_ids[:token_len]
return scheduler, tracker

Expand All @@ -219,10 +267,10 @@ def test_cached_request_with_spec_decode_does_not_save_scheduled_drafts():
)

assert meta.requests == []
assert scheduler._test_scheduler_context_calls == []
tracker = scheduler._request_trackers["req-0"]
assert tracker.token_len == 44
assert tracker.num_saved_tokens == 32
assert tracker.allocated_block_ids == ([0, 1, 2],)


def test_cached_request_without_spec_decode_keeps_current_step_save_overlap():
Expand Down Expand Up @@ -259,7 +307,6 @@ def test_decode_tracking_is_skipped_by_default(kv_role):
assert meta.requests == []
assert tracker.token_len == 47
assert tracker.num_saved_tokens == 32
assert tracker.allocated_block_ids == ([0, 1, 2],)
assert tracker.token_ids == list(range(47))


Expand All @@ -275,7 +322,7 @@ def test_fresh_consumer_first_decode_save_can_backfill_missing_prompt():
request.block_hashes = [b"h0", b"h1", b"h2"]
request.all_token_ids = list(range(48))
request.num_output_placeholders = 0
scheduler._unfinished_requests["req-0"] = (request, request.block_ids)
_set_unfinished_request(scheduler, request, request.block_ids)

# The consumer does not save during prefill, so its first decode save must
# still cover the prompt. The worker deduplicates prompt blocks already in
Expand Down Expand Up @@ -348,7 +395,6 @@ def test_preemption_resets_tracker():

tracker = scheduler._request_trackers["req-0"]
assert tracker.token_len == 0
assert tracker.allocated_block_ids == ()
assert tracker.num_saved_tokens == 0
assert tracker.token_ids is None
assert tracker.has_pending_offload is False
Expand Down Expand Up @@ -388,7 +434,7 @@ def _make_pending_load_unfinished_request(
block_hashes=block_hashes,
num_output_placeholders=0,
)
scheduler._unfinished_requests["req-0"] = (request, block_ids)
_set_unfinished_request(scheduler, request, block_ids)


def _make_pending_load_scheduler_output() -> SimpleNamespace:
Expand Down Expand Up @@ -437,6 +483,8 @@ def test_pending_load_does_not_co_queue_save():
# Load is still issued as planned.
assert req_meta.load_spec is not None
assert req_meta.load_spec.can_load is True
assert req_meta.block_ids == ([0, 1, 2],)
assert scheduler._test_scheduler_context_calls == ["req-0"]
# And the save watermark does not advance for a save that was never queued.
tracker = scheduler._request_trackers["req-0"]
assert tracker.num_saved_tokens == 0
Expand All @@ -455,7 +503,7 @@ def _make_resumed_unfinished_request(
num_computed_tokens=num_computed_tokens,
num_output_placeholders=0,
)
scheduler._unfinished_requests["req-0"] = (request, ([0, 1],))
_set_unfinished_request(scheduler, request, ([0, 1, 2],))


def _make_resumed_scheduler_output(*, num_scheduled_tokens: int) -> SimpleNamespace:
Expand Down Expand Up @@ -504,6 +552,7 @@ def test_resumed_from_preemption_with_load_skips_save():
assert req_meta.can_save is False
assert req_meta.load_spec is not None
assert req_meta.load_spec.can_load is True
assert req_meta.block_ids == ([0, 1, 2],)
tracker = scheduler._request_trackers["req-0"]
assert tracker.num_saved_tokens == 0

Expand Down Expand Up @@ -531,17 +580,8 @@ def test_resumed_from_preemption_without_load_still_saves():
assert tracker.num_saved_tokens == 48


def test_running_request_not_in_resumed_req_ids_appends_blocks():
"""Regression: the replace-vs-append choice must follow the scheduler's
cached_reqs.resumed_req_ids, NOT connector-local preemption history.

A running request that is not resumed this step carries a *delta*
new_block_ids and must be APPENDED to the tracker's existing blocks.
Treating it as resumed would replace allocated_block_ids with just the
delta while token_len stays at the full computed length, so the store
path's block_ids[start // block_size] runs off the end (the
"list index out of range" / token_len >> len(block_ids) bug).
"""
def test_running_request_uses_authoritative_block_table():
"""A running request's delta block update is not used as worker metadata."""
scheduler = _make_bare_scheduler()
_add_unfinished_request(
scheduler,
Expand All @@ -555,45 +595,34 @@ def test_running_request_not_in_resumed_req_ids_appends_blocks():

meta = scheduler.build_connector_meta(out)

tracker = scheduler._request_trackers["req-0"]
# Delta [2] appended to existing [0, 1] (decode path), not replaced by [2].
assert tracker.allocated_block_ids == ([0, 1, 2],)
# token_len stays covered by the block table: no store-path under-count.
blocks_held = sum(len(g) for g in tracker.allocated_block_ids)
assert tracker.token_len // scheduler._block_size <= blocks_held
assert len(meta.requests) == 1
assert meta.requests[0].token_len_chunk == 48
req_meta = meta.requests[0]
assert req_meta.token_len_chunk == 48
assert req_meta.block_ids == ([0, 1, 2],)
assert scheduler._test_scheduler_context_calls == ["req-0"]


def test_resumed_request_in_resumed_req_ids_replaces_blocks():
"""A request the scheduler marks resumed gets the FULL block table in
new_block_ids and must REPLACE the tracker's blocks (not append), even if
a stale tracker from before preemption is still present."""
def test_resumed_request_reinitializes_tracker_and_uses_current_blocks():
scheduler = _make_bare_scheduler()
_make_resumed_unfinished_request(
scheduler,
token_ids=list(range(48)),
block_hashes=[b"h0", b"h1", b"h2"],
num_computed_tokens=0,
)
# Stale pre-preemption tracker that must be overwritten, not appended to.
scheduler._request_trackers["req-0"] = RequestTracker(
req_id="req-0",
token_len=99,
allocated_block_ids=([7, 8, 9],),
num_saved_tokens=0,
)

scheduler.build_connector_meta(
meta = scheduler.build_connector_meta(
_make_resumed_scheduler_output(num_scheduled_tokens=48)
)

tracker = scheduler._request_trackers["req-0"]
# Replaced with the full table from new_block_ids, not appended to [7,8,9].
assert tracker.allocated_block_ids == ([0, 1, 2],)
assert tracker.token_len == 48
blocks_held = sum(len(g) for g in tracker.allocated_block_ids)
assert tracker.token_len // scheduler._block_size <= blocks_held
assert meta.requests[0].block_ids == ([0, 1, 2],)


# Focused tests for ReqMeta.from_request_tracker — the centralized guard that
Expand All @@ -607,7 +636,6 @@ def test_from_request_tracker_load_overrides_caller_skip_save():
tracker = RequestTracker(
req_id="req-0",
token_len=48,
allocated_block_ids=([0, 1, 2],),
num_saved_tokens=0,
)
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=48, can_load=True)
Expand All @@ -632,7 +660,6 @@ def test_from_request_tracker_load_with_can_load_false_still_saves():
tracker = RequestTracker(
req_id="req-0",
token_len=48,
allocated_block_ids=([0, 1, 2],),
num_saved_tokens=0,
)
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=48, can_load=False)
Expand All @@ -656,7 +683,6 @@ def test_from_request_tracker_no_load_saves_normally():
tracker = RequestTracker(
req_id="req-0",
token_len=48,
allocated_block_ids=([0, 1, 2],),
num_saved_tokens=0,
)

Expand Down Expand Up @@ -876,11 +902,10 @@ def _add_pending_partial_tail_request(
num_output_placeholders=0,
num_prompt_tokens=12,
)
scheduler._unfinished_requests["req-0"] = (request, block_ids)
_set_unfinished_request(scheduler, request, block_ids)
scheduler._request_trackers["req-0"] = RequestTracker(
req_id="req-0",
token_len=num_tokens,
allocated_block_ids=block_ids,
num_saved_tokens=0,
token_ids=list(range(num_tokens)),
prefill_end_tokens=num_tokens,
Expand Down Expand Up @@ -959,11 +984,10 @@ def test_resumed_partial_tail_attached_to_save_keeps_handoff_boundary():
num_output_placeholders=0,
num_prompt_tokens=36,
)
scheduler._unfinished_requests["req-0"] = (request, ([0, 1],))
_set_unfinished_request(scheduler, request, ([0, 1, 2],))
scheduler._request_trackers["req-0"] = RequestTracker(
req_id="req-0",
token_len=44,
allocated_block_ids=([0, 1],),
num_saved_tokens=32,
token_ids=list(range(44)),
prefill_end_tokens=48,
Expand Down Expand Up @@ -994,18 +1018,44 @@ def test_partial_tail_cow_block_is_referenced_for_the_job():
block_hashes=[b"h0", b"h1", b"h2"],
block_ids=([0],),
)
pool = scheduler._gpu_block_pool
pool = _get_gpu_block_pool(scheduler)

meta = scheduler.build_connector_meta(out)

store_job_id = meta.requests[0].store_job_id
# It leads the list, as in `pop_blocks_for_free`, so that the reversed free
# puts it last in eviction priority.
assert scheduler._pinned_saves[store_job_id][0] == [7, 0]
assert meta.requests[0].block_ids == ([NULL_BLOCK_ID],)
assert scheduler._pinned_saves[store_job_id][0] == [7]
assert pool.blocks[NULL_BLOCK_ID].ref_cnt == 0
assert pool.blocks[7].ref_cnt == 1

scheduler.update_connector_output(_make_worker_output({store_job_id: 1}))
assert pool.blocks[7].ref_cnt == 0


def test_partial_tail_and_current_blocks_are_pinned_once():
scheduler = _make_bare_scheduler(hash_block_size=4, enable_partial_hash_hits=True)
out = _add_pending_partial_tail_request(
scheduler,
num_tokens=12,
block_hashes=[b"h0", b"h1", b"h2"],
block_ids=([7, NULL_BLOCK_ID], [7, 8]),
)
pool = _get_gpu_block_pool(scheduler)

meta = scheduler.build_connector_meta(out)

[req_meta] = meta.requests
assert req_meta.block_ids == ([7, NULL_BLOCK_ID], [7, 8])
store_job_id = req_meta.store_job_id
assert scheduler._pinned_saves[store_job_id][0] == [7, 8]
assert pool.blocks[NULL_BLOCK_ID].ref_cnt == 0
assert pool.blocks[7].ref_cnt == 1
assert pool.blocks[8].ref_cnt == 1

scheduler.update_connector_output(_make_worker_output({store_job_id: 1}))

assert pool.blocks[7].ref_cnt == 0
assert pool.blocks[8].ref_cnt == 0


def test_store_job_blocks_are_released_once_every_rank_reports():
Expand All @@ -1021,7 +1071,7 @@ def test_store_job_blocks_are_released_once_every_rank_reports():
block_hashes=[b"h0", b"h1", b"h2"],
prefill_end_tokens=48,
)
pool = scheduler._gpu_block_pool
pool = _get_gpu_block_pool(scheduler)
assert scheduler.has_pending_push_work() is False

meta = scheduler.build_connector_meta(
Expand Down
Loading
Loading