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
76 changes: 71 additions & 5 deletions tests/v1/kv_offload/tiering/test_async_lookup.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
import pytest

from vllm.v1.kv_offload.base import OffloadKey, ReqContext, make_offload_key
from vllm.v1.kv_offload.tiering.async_lookup import AsyncLookupManager
from vllm.v1.kv_offload.tiering.async_lookup import AsyncLookupManager, LookupPhase


def _key(i: int) -> OffloadKey:
Expand Down Expand Up @@ -46,10 +46,13 @@ def test_new_key_returns_none(self):
def test_found_key_returns_true(self):
mgr = InMemoryLookupManager(existing_keys={_key(1)})
assert mgr.lookup(_key(1), _ctx()) is None
assert mgr._lookup_state[_key(1)].phase is LookupPhase.PENDING
mgr.flush()
assert mgr._lookup_state[_key(1)].phase is LookupPhase.IN_FLIGHT
mgr._results_ready.wait()
mgr._results_ready.clear()
assert mgr.lookup(_key(1), _ctx()) is True
assert mgr._lookup_state[_key(1)].phase is LookupPhase.RESOLVED
mgr.shutdown()

def test_not_found_key_returns_false(self):
Expand Down Expand Up @@ -106,6 +109,68 @@ def test_cleanup_preserves_shared_entries(self):
assert _key(1) not in mgr._lookup_state
mgr.shutdown()

def test_cleanup_reuses_in_flight_probe(self, monkeypatch: pytest.MonkeyPatch):
"""A replacement request shares the in-flight probe and its verdict."""
key = _key(1)
mgr = InMemoryLookupManager(existing_keys={key})
ctx_b = _ctx("req_b")
probe_started = threading.Event()
release_probe = threading.Event()
batch_lookup = mgr.batch_lookup

def blocking_lookup(keys, req_context):
probe_started.set()
if not release_probe.wait(timeout=5):
raise TimeoutError("Test did not release the backend probe")
return batch_lookup(keys, req_context)

monkeypatch.setattr(mgr, "batch_lookup", blocking_lookup)
try:
assert mgr.lookup(key, _ctx("req_a")) is None
mgr.flush()
assert probe_started.wait(timeout=5)
assert mgr._lookup_state[key].phase is LookupPhase.IN_FLIGHT
mgr.cleanup("req_a")
assert mgr.lookup(key, ctx_b) is None
assert mgr._lookup_state[key].phase is LookupPhase.IN_FLIGHT
mgr.flush()

release_probe.set()
batch = mgr._pending_results.get(timeout=5)
mgr._pending_results.put(batch)
replacement_result = mgr.lookup(key, ctx_b)
finally:
release_probe.set()
mgr.shutdown()

assert mgr.batch_lookup_calls == 1
assert replacement_result is True

@pytest.mark.parametrize(
"reclaim_at_shutdown", [False, True], ids=["flush", "shutdown"]
)
def test_unclaimed_probe_reclaimed_without_lookup(self, reclaim_at_shutdown: bool):
"""A completed orphan is released by flush or shutdown without a lookup."""
key = _key(1)
mgr = InMemoryLookupManager(existing_keys={key})
try:
mgr.lookup(key, _ctx("req_a"))
mgr.flush()
mgr.cleanup("req_a")
assert key in mgr._lookup_state
assert not mgr._req_keys

if not reclaim_at_shutdown:
batch = mgr._pending_results.get(timeout=5)
mgr._pending_results.put(batch)
mgr.flush()
assert key not in mgr._lookup_state
finally:
mgr.shutdown()

assert not mgr._lookup_state
assert mgr._pending_results.empty()

def test_stale_result_ignored_after_cleanup_and_key_reuse(self):
key = _key(1)
mgr = InMemoryLookupManager()
Expand All @@ -119,11 +184,12 @@ def test_stale_result_ignored_after_cleanup_and_key_reuse(self):
generation = mgr._lookup_state[key].generation
assert generation != stale_generation

mgr.flush()
mgr._results_ready.wait()
mgr._results_ready.clear()
current_result = mgr._pending_results.get(timeout=5)
mgr._pending_results.put([(key, stale_generation, True)])
mgr.drain_results()
assert mgr.lookup(key, ctx_b) is None

mgr._pending_results.put([(key, generation, False)])
mgr._pending_results.put(current_result)
mgr.drain_results()
assert mgr.lookup(key, ctx_b) is False
mgr.shutdown()
Expand Down
57 changes: 41 additions & 16 deletions vllm/v1/kv_offload/tiering/async_lookup.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,26 +27,40 @@
flush() is called once per step from the tier's on_schedule_end(), posting
the entire batch as a single queue item so the background thread sees one
batch per step.
drain_results() is called before any lookup() calls in the same step, so
lookup() is a pure OrderedDict operation.
Results are drained on the first lookup after each flush, at flush(), and
after worker shutdown. In-flight lookups with no remaining request references
are retained until their results are drained, allowing new requests to share
the same probe.
"""

import queue
import threading
from abc import ABC, abstractmethod
from collections.abc import Collection, Iterable
from dataclasses import dataclass, field
from enum import Enum, auto

from vllm.logger import init_logger
from vllm.v1.kv_offload.base import OffloadKey, ReqContext

logger = init_logger(__name__)


class LookupPhase(Enum):
"""Lifecycle phase of a lookup probe."""

PENDING = auto() # Accumulated in _lookup_batch, but not yet submitted.
IN_FLIGHT = auto() # Submitted to the worker, but not yet resolved.
RESOLVED = auto() # A final verdict has been applied to the state.


@dataclass(slots=True)
class LookupState:
generation: int
result: bool | None = None # True (found), False (not found), None
phase: LookupPhase = LookupPhase.PENDING
# None while pending/in flight; True if the key exists; False if absent or
# explicitly marked missing after a failed load.
result: bool | None = None
request_ids: set[str] = field(default_factory=set) # requests asking for the lookup


Expand Down Expand Up @@ -155,24 +169,28 @@ def flush(self) -> None:
Called once per step from on_schedule_end() after all lookup() calls
are done. The worker receives the full batch and processes it during
the model-execution window, maximising time available before the next
step's drain_results(). Safe to call with an empty batch (no-op).
step's drain_results(). Also drains completed lookups when there
are no new keys to submit.
"""
self.drain_results()
self._need_to_drain = True
batch = self._lookup_batch
self._lookup_batch = []
batch = [
(key, req_context, generation)
for key, req_context, generation in batch
if (state := self._lookup_state.get(key)) is not None
and state.generation == generation
]
if batch:
self._lookup_queue.put(batch)
in_flight_batch = []
for key, req_context, generation in batch:
state = self._lookup_state.get(key)
if state is None or state.generation != generation:
continue
assert state.phase is LookupPhase.PENDING
state.phase = LookupPhase.IN_FLIGHT
in_flight_batch.append((key, req_context, generation))
if in_flight_batch:
self._lookup_queue.put(in_flight_batch)

def drain_results(self) -> None:
"""Apply pending worker results to _lookup_state.

Called from lookup() before checking state.
Called from lookup(), flush(), and shutdown() on the scheduler thread.
"""
while True:
try:
Expand All @@ -183,6 +201,10 @@ def drain_results(self) -> None:
state = self._lookup_state.get(key)
if state is None or state.generation != generation:
continue
if not state.request_ids:
del self._lookup_state[key]
continue
assert state.phase is LookupPhase.IN_FLIGHT
# Each lookup generation is enqueued exactly once. A matching
# generation must not receive a second result; stale
# generations were discarded above.
Expand All @@ -192,6 +214,7 @@ def drain_results(self) -> None:
"failed-load livelock"
)
state.result = result
state.phase = LookupPhase.RESOLVED

def mark_miss(self, keys: Collection[OffloadKey]) -> None:
"""Force the cached verdict for ``keys`` to False after a failed load, so
Expand All @@ -201,23 +224,25 @@ def mark_miss(self, keys: Collection[OffloadKey]) -> None:
state = self._lookup_state.get(key)
if state is not None:
state.result = False
state.phase = LookupPhase.RESOLVED

def cleanup(self, req_id: str) -> None:
"""Remove entries no longer needed by any active request.
"""Release request references, retaining in-flight lookups.

Called from the tier's on_request_finished(). Uses the reverse
index to visit only keys associated with this request.
"""
for key in self._req_keys.pop(req_id, ()):
state = self._lookup_state[key]
state.request_ids.discard(req_id)
if not state.request_ids:
if not state.request_ids and state.phase is not LookupPhase.IN_FLIGHT:
del self._lookup_state[key]

def shutdown(self) -> None:
"""Stop the worker thread."""
"""Stop the worker thread and drain completed lookups."""
self._lookup_queue.put(None) # unblock _worker from _lookup_queue.get()
self._thread.join()
self.drain_results()

# ------------------------------------------------------------------
# Internal helpers
Expand Down
Loading