Skip to content
1 change: 1 addition & 0 deletions docs/features/mooncake_store_connector_usage.md
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,7 @@ the vLLM JSON config.
### kv_connector_extra_config

- `load_async` (bool): Enable asynchronous loading for better compute-I/O overlap. Default: `true`.
- `lookup_async` (bool): Run the external prefix-cache lookup on a background thread so it never blocks the scheduler step. The request is held until the in-flight lookup completes, then resumed on a later step. Default: `false`.
- `enable_cross_layers_blocks` (bool): Enable cross-layer block packing for reduced store operations. Default: `false`.
- `lookup_rpc_port` (int): Custom port for the ZMQ lookup RPC socket. Default: `0`.
- `cache_prefix` (str): Namespace prepended to every store key. Lets separate deployments share one Mooncake master without polluting each other — instances configured with different prefixes never see each other's cached blocks, even for identical prompts. All instances that should share a prefix cache must use the same value. Default: `""` (no prefix; keys are byte-identical to the unprefixed format).
Expand Down
127 changes: 126 additions & 1 deletion tests/v1/kv_connector/unit/test_mooncake_store_connector.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import threading
import time
from unittest.mock import MagicMock, patch

from vllm.config import set_current_vllm_config
Expand Down Expand Up @@ -406,7 +408,9 @@ def test_lookup_key_client_lookup_prepends_typed_tag():
fake_socket = mock_make_socket.return_value
fake_socket.recv.return_value = (5).to_bytes(4, "big")

assert client.lookup(token_len=128, block_hashes=[]) == 5
# Blocking lookup (non_block defaults to False) runs on the executor and
# returns the resolved hit length.
assert client.lookup("req0", token_len=128, block_hashes=[]) == 5

sent_frames = fake_socket.send_multipart.call_args[0][0]
assert sent_frames[0] == protocol.LOOKUP_MSG
Expand Down Expand Up @@ -435,6 +439,127 @@ def test_lookup_key_client_reset_uses_typed_protocol():
assert client.reset() is False


def _poll_lookup(client, req_id, token_len=128, block_hashes=(), timeout=5.0):
"""Drive non-blocking lookup until the executor completes it."""
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
result = client.lookup(req_id, token_len, list(block_hashes), non_block=True)
if result is not None:
return result
time.sleep(0.005)
return None


def _gated_recv(gate: threading.Event, value: int):
"""Mock recv side-effect that blocks until ``gate`` is set, so the
executor's lookup can be held pending deterministically."""

def recv():
gate.wait()
return value.to_bytes(4, "big")

return recv


def test_lookup_key_client_non_block_lookup_async():
"""Non-blocking lookup defers to the executor: None first, hit once the
Future resolves."""
vllm_config = _make_vllm_config()

with patch(
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
"worker.make_zmq_socket"
) as mock_make_socket:
client = worker.LookupKeyClient(vllm_config)

fake_socket = mock_make_socket.return_value
# Hold the executor's lookup pending until we release the gate.
gate = threading.Event()
fake_socket.recv.side_effect = _gated_recv(gate, 7)

# First query submits the lookup and returns None while it is in flight.
assert client.lookup("req1", 128, [], non_block=True) is None
# Release the executor; a later poll returns the hit length.
gate.set()
assert _poll_lookup(client, "req1") == 7
# Future is consumed (popped) on read.
assert "req1" not in client.futures


def test_lookup_key_client_discard_clears_state():
"""discard() drops a completed lookup Future so it is not served stale."""
vllm_config = _make_vllm_config()

with patch(
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
"worker.make_zmq_socket"
) as mock_make_socket:
client = worker.LookupKeyClient(vllm_config)

fake_socket = mock_make_socket.return_value
gate = threading.Event()
fake_socket.recv.side_effect = _gated_recv(gate, 9)

# Submit while gated so the call returns None and the Future stays in
# `futures` (unconsumed) once it resolves.
assert client.lookup("req2", 128, [], non_block=True) is None
gate.set()
deadline = time.monotonic() + 5.0
while time.monotonic() < deadline:
if client.futures["req2"].done():
break
time.sleep(0.005)
# discard() drops the completed result before any lookup consumes it.
client.discard("req2")
assert "req2" not in client.futures
# A fresh query re-submits rather than returning a stale value: hold the
# gate so the resubmitted lookup stays in flight.
gate.clear()
assert client.lookup("req2", 128, [], non_block=True) is None
gate.set() # release the executor so the worker thread can drain


def test_get_num_new_matched_tokens_async_defers_then_reports():
"""Async lookup returns (None, False) until ready, then the hit count."""
vllm_config = create_vllm_config(
kv_connector="MooncakeStoreConnector",
kv_role="kv_both",
kv_connector_extra_config={"lookup_async": True},
)
kv_cache_config = _make_kv_cache_config()

with (
set_current_vllm_config(vllm_config),
patch(
"vllm.distributed.kv_transfer.kv_connector.v1.mooncake.store."
"scheduler.LookupKeyClient"
) as mock_client_cls,
):
sched = scheduler.MooncakeStoreScheduler(vllm_config, kv_cache_config)

assert sched.lookup_async is True
mock_client = mock_client_cls.return_value

block_size = sched._block_size
request = MagicMock()
request.request_id = "r1"
request.num_tokens = 4 * block_size
request.block_hashes = []

# Lookup not ready -> defer.
mock_client.lookup.return_value = None
assert sched.get_num_new_matched_tokens(request, 0) == (None, False)
assert "r1" not in sched.load_specs

# Lookup ready with a hit -> report need_to_allocate + async-load flag.
hit = 3 * block_size
mock_client.lookup.return_value = hit
need, load_async = sched.get_num_new_matched_tokens(request, 0)
assert need == hit
assert load_async == sched.load_async
assert sched.load_specs["r1"].kvpool_cached_tokens == hit


def test_protocol_tags_are_distinct_and_non_empty():
"""Protocol tags must be unique and non-empty to avoid collision."""
tags = {protocol.LOOKUP_MSG, protocol.RESET_MSG}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
def _make_bare_scheduler() -> MooncakeStoreScheduler:
scheduler = object.__new__(MooncakeStoreScheduler)
scheduler.kv_role = "kv_both"
scheduler.lookup_async = False
scheduler._block_size = 16
scheduler.load_specs = {}
scheduler._preempted_req_ids = set()
Expand Down Expand Up @@ -405,7 +406,13 @@ class _StubLookupClient:
def __init__(self, hit_tokens: int) -> None:
self._hit_tokens = hit_tokens

def lookup(self, token_len: int, block_hashes: list[bytes]) -> int:
def lookup(
self,
req_id: str,
token_len: int,
block_hashes: list[bytes],
non_block: bool = False,
) -> int:
return self._hit_tokens


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -176,7 +176,7 @@ def get_num_new_matched_tokens(
self,
request: Request,
num_computed_tokens: int,
) -> tuple[int, bool]:
) -> tuple[int | None, bool]:
assert self.connector_scheduler is not None
return self.connector_scheduler.get_num_new_matched_tokens(
request, num_computed_tokens
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -54,9 +54,9 @@ def __init__(
):
assert vllm_config.kv_transfer_config is not None
self.kv_role = vllm_config.kv_transfer_config.kv_role
self.load_async = vllm_config.kv_transfer_config.kv_connector_extra_config.get(
"load_async", True
)
kvc_extra_config = vllm_config.kv_transfer_config.kv_connector_extra_config
self.load_async = kvc_extra_config.get("load_async", True)
self.lookup_async = kvc_extra_config.get("lookup_async", False)
self.client = LookupKeyClient(vllm_config)

# Align with the engine's own scheduler_block_size and hash_block_size.
Expand All @@ -75,14 +75,26 @@ def get_num_new_matched_tokens(
self,
request: Request,
num_computed_tokens: int,
) -> tuple[int, bool]:
"""Check for external KV cache hit."""
) -> tuple[int | None, bool]:
"""Check for external KV cache hit.

Returns ``(None, False)`` when an async lookup is still in flight,
signaling the scheduler to retry this request on a later step.
"""
# Look up against the full prefill range, not just the prompt.
token_len = request.num_tokens // self._block_size * self._block_size
if token_len < self._block_size:
return 0, False

num_external_hit_tokens = self.client.lookup(token_len, request.block_hashes)
num_external_hit_tokens = self.client.lookup(
request.request_id,
token_len,
request.block_hashes,
non_block=self.lookup_async,
)
if num_external_hit_tokens is None:
# Lookup not ready yet; scheduler will retry on a later step.
return None, False

if num_external_hit_tokens == request.num_tokens:
# Leave a sub-block tail uncomputed for sampling, on a block
Expand Down Expand Up @@ -158,6 +170,7 @@ def build_connector_meta(
force_skip_save = self.kv_role == "kv_consumer"

for finished_req_id in scheduler_output.finished_req_ids:
self.client.discard(finished_req_id)
self.load_specs.pop(finished_req_id, None)
self._request_trackers.pop(finished_req_id, None)
self._unfinished_requests.pop(finished_req_id, None)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import time
from collections import defaultdict
from collections.abc import Callable
from concurrent.futures import Future, ThreadPoolExecutor
from dataclasses import dataclass
from typing import Any, Literal, TypeVar

Expand Down Expand Up @@ -1560,7 +1561,13 @@ def __init__(self, vllm_config: VllmConfig):
bind=False,
)

def lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
# Async lookup support
self.executor = ThreadPoolExecutor(
max_workers=1, thread_name_prefix="MooncakeLookupClient"
)
self.futures: dict[str, Future[int]] = {}

def _lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
hash_strs = [h.hex() for h in block_hashes]
hash_frames = self.encoder.encode(hash_strs)
token_len_bytes = token_len.to_bytes(4, byteorder="big")
Expand All @@ -1570,7 +1577,36 @@ def lookup(self, token_len: int, block_hashes: list[BlockHash]) -> int:
result = int.from_bytes(resp, "big")
return result

def reset(self) -> bool:
def lookup(
self,
req_id: str,
token_len: int,
block_hashes: list[BlockHash],
non_block: bool = False,
) -> int | None:
"""If non_block is True, will return None until the result is ready,
so the caller retries on a later step."""
future = self.futures.get(req_id)
if future is None:
future = self.executor.submit(self._lookup, token_len, list(block_hashes))
self.futures[req_id] = future
if non_block and not future.done():
return None
try:
return future.result()
except Exception as e:
logger.error("Async Mooncake lookup failed for %s: %s", req_id, e)
return 0
finally:
del self.futures[req_id]

def discard(self, req_id: str) -> None:
"""Drop any cached/in-flight lookup for ``req_id`` (e.g. on abort)."""
future = self.futures.pop(req_id, None)
if future is not None:
future.cancel()

def _reset(self) -> bool:
"""Trigger ``store.remove_all(force=True)`` on worker rank 0.

Ordering assumption: caller MUST ensure no in-flight Mooncake
Expand All @@ -1582,7 +1618,11 @@ def reset(self) -> bool:
resp = self.socket.recv()
return bytes(resp) == RESP_OK

def reset(self) -> bool:
return self.executor.submit(self._reset).result()

def close(self):
self.executor.shutdown(wait=False, cancel_futures=True)
self.socket.close(linger=0)


Expand Down
Loading