diff --git a/tests/v1/kv_offload/tiering/test_async_lookup.py b/tests/v1/kv_offload/tiering/test_async_lookup.py index ff5cbcd5dd44..3649c8df0488 100644 --- a/tests/v1/kv_offload/tiering/test_async_lookup.py +++ b/tests/v1/kv_offload/tiering/test_async_lookup.py @@ -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: @@ -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): @@ -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() @@ -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() diff --git a/vllm/v1/kv_offload/tiering/async_lookup.py b/vllm/v1/kv_offload/tiering/async_lookup.py index f111f233b613..c08c63d01875 100644 --- a/vllm/v1/kv_offload/tiering/async_lookup.py +++ b/vllm/v1/kv_offload/tiering/async_lookup.py @@ -27,8 +27,10 @@ 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 @@ -36,6 +38,7 @@ 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 @@ -43,10 +46,21 @@ 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 @@ -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: @@ -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. @@ -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 @@ -201,9 +224,10 @@ 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. @@ -211,13 +235,14 @@ def cleanup(self, req_id: str) -> None: 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