diff --git a/lmcache/v1/distributed/storage_controllers/prefetch_controller.py b/lmcache/v1/distributed/storage_controllers/prefetch_controller.py index 54451854ca2..0e6f8e25355 100644 --- a/lmcache/v1/distributed/storage_controllers/prefetch_controller.py +++ b/lmcache/v1/distributed/storage_controllers/prefetch_controller.py @@ -193,6 +193,18 @@ class PrefetchPhase(enum.Enum): PLAN_AND_LOAD = enum.auto() +@dataclass(frozen=True) +class PrefetchResult: + """Retained keys and whether reserving their load buffers failed. + + Allocation failure or contention does not prove stored pages are absent. + ``found`` indexes the keys submitted to the prefetch request. + """ + + found: Bitmap + reservation_failed: bool = False + + @dataclass class InFlightPrefetchRequest: """Tracks a single prefetch request across its lifecycle phases.""" @@ -245,6 +257,8 @@ class InFlightPrefetchRequest: """Maps object_group_id to that group's layout (one ``MemoryLayoutDesc`` describes a single group's MemoryObj). Covers every object group.""" + reservation_failed: bool = False + def all_lookups_done(self) -> bool: return len(self.pending_lookup_tasks) == 0 @@ -338,6 +352,7 @@ def __init__( self._prefetch_results_lock = threading.Lock() self._prefetch_results_cv = threading.Condition(self._prefetch_results_lock) self._completed_results: dict[PrefetchRequestId, Bitmap] = {} + self._completed_reservation_failures: set[PrefetchRequestId] = set() # Map eventfds to adapter indices for quick lookup in poll. # Relies on the L2AdapterInterface contract that every adapter @@ -488,12 +503,26 @@ def query_prefetch_result(self, request_id: PrefetchRequestId) -> Bitmap | None: query_lookup_result after calling this function, otherwise it will get None forever. """ + result = self.query_prefetch_result_detailed(request_id) + return result.found if result is not None else None + + def query_prefetch_result_detailed( + self, request_id: PrefetchRequestId + ) -> PrefetchResult | None: + """Consume hits and allocation evidence for ``request_id`` atomically. + + Returns None while pending or already consumed. Bitmap-only queries + consume the same result and release the same lookup bookkeeping. + """ with self._prefetch_results_lock: - result = self._completed_results.pop(request_id, None) - if result is not None: - with self._lookup_results_lock: - self._completed_lookups.pop(request_id, None) - return result + found = self._completed_results.pop(request_id, None) + if found is None: + return None + failed = request_id in self._completed_reservation_failures + self._completed_reservation_failures.discard(request_id) + with self._lookup_results_lock: + self._completed_lookups.pop(request_id, None) + return PrefetchResult(found, failed) def wait_prefetch_result( self, request_id: PrefetchRequestId, timeout: float @@ -1102,6 +1131,7 @@ def _reserve_load_buffers( request.write_reserved_objs[key] = mem_obj reserved.add(key) continue + request.reservation_failed = True if err == L1Error.OUT_OF_MEMORY: oom_keys.append(key) elif err == L1Error.KEY_NOT_WRITABLE: @@ -1600,6 +1630,9 @@ def _release_l2_locks( def _complete_request(self, request_id: PrefetchRequestId, result: Bitmap) -> None: """Store the retained-key bitmap and remove from in-flight tracking.""" with self._prefetch_results_lock: + request = self._in_flight_requests.get(request_id) + if request is not None and request.reservation_failed: + self._completed_reservation_failures.add(request_id) self._completed_results[request_id] = result # Wake any WAIT_PREFETCH_STATUS handler blocked on this result. self._prefetch_results_cv.notify_all() diff --git a/lmcache/v1/distributed/storage_manager.py b/lmcache/v1/distributed/storage_manager.py index e8dba946e1f..72a943ac218 100644 --- a/lmcache/v1/distributed/storage_manager.py +++ b/lmcache/v1/distributed/storage_manager.py @@ -56,6 +56,9 @@ PrefetchController, StoreController, ) +from lmcache.v1.distributed.storage_controllers.prefetch_controller import ( + PrefetchResult, +) from lmcache.v1.distributed.storage_controllers.prefetch_policy import ( create_prefetch_policy, ) @@ -1030,13 +1033,27 @@ def query_prefetch_status( done, None if it's still in progress. Derive the prefix hit count via ``count_leading_ones``. """ + result = self.query_prefetch_status_detailed(handle) + return result.found if result is not None else None + + def query_prefetch_status_detailed( + self, handle: PrefetchHandle + ) -> PrefetchResult | None: + """Consume hits and load-allocation evidence for ``handle``. + + Returns None while pending or already consumed. Found bits index the + original requested keys. Bitmap-only queries consume the same result. + """ l2_r: Bitmap | None = None + reservation_failed = False if handle.prefetch_request_id != -1: - l2_r = self._prefetch_controller.query_prefetch_result( + result = self._prefetch_controller.query_prefetch_result_detailed( handle.prefetch_request_id ) - if l2_r is None: + if result is None: return None + l2_r = result.found + reservation_failed = result.reservation_failed found = self._combine_found(handle, l2_r) # popcount (not count_leading_ones) so the log is accurate for @@ -1060,7 +1077,7 @@ def query_prefetch_status( handle.external_request_id, handle.prefetch_request_id, ) - return found + return PrefetchResult(found, reservation_failed) def touch_l1_keys(self, keys: list[ObjectKey]): """ diff --git a/lmcache/v1/multiprocess/checkpoint_storage.py b/lmcache/v1/multiprocess/checkpoint_storage.py index 9918aeb8562..65a16896a35 100644 --- a/lmcache/v1/multiprocess/checkpoint_storage.py +++ b/lmcache/v1/multiprocess/checkpoint_storage.py @@ -606,8 +606,8 @@ def poll_retrieve(self, lease_id: str) -> CheckpointSlots | bool | None: The storage manager loads pages from L2 only into RAM it reserved for all of them, so a full RAM leaves stored pages unread. A lookup with unread pages is therefore repeated once RAM has room for them, - and only a repeated lookup that had that room invalidates the - generation. The lease stays pending while eviction makes room, for + and only a repeated lookup without allocation failure invalidates + the generation. The lease stays pending while eviction makes room, for at most the storage admission timeout; a lease that never gets room misses without invalidating its generation. @@ -621,13 +621,16 @@ def poll_retrieve(self, lease_id: str) -> CheckpointSlots | bool | None: return lease.slots if lease.handle is None: return self._repeat_with_room(lease_id, lease) - found = self._storage.query_prefetch_status(lease.handle) - if found is None: + result = self._storage.query_prefetch_status_detailed(lease.handle) + if result is None: return None + found = result.found readable_keys = [key for i, key in enumerate(lease.keys) if found.test(i)] if len(readable_keys) != len(lease.keys) or lease.cancelled: self._storage.finish_read_prefetched(readable_keys) - return self._repeat_or_miss(lease_id, lease, readable_keys) + return self._repeat_or_miss( + lease_id, lease, readable_keys, result.reservation_failed + ) try: keys, objects = self._storage.unsafe_read(lease.keys) if keys != lease.keys or len(objects) != len(keys): @@ -652,30 +655,44 @@ def poll_retrieve(self, lease_id: str) -> CheckpointSlots | bool | None: raise def _repeat_or_miss( - self, lease_id: str, lease: _RetrieveLease, readable_keys: list[ObjectKey] + self, + lease_id: str, + lease: _RetrieveLease, + readable_keys: list[ObjectKey], + reservation_failed: bool, ) -> bool | None: """Repeat or end a lookup whose pages were not all pinned. Called with the lookup's read locks already released. Returns None when the lookup will be repeated, otherwise False after forgetting the - lease. Only a repeated lookup that had room in RAM invalidates. + lease. Only a repeated lookup without allocation failure invalidates. """ lease.readable = len(readable_keys) if lease.cancelled: return self._miss(lease_id, lease, "was cancelled by the engine") groups = checkpoint_page_groups(lease.manifest) readable = set(readable_keys) + alignment = self._storage.l1_memory_desc.align_bytes lease.unread_bytes = sum( - groups[key.object_group_id].page_bytes + ((groups[key.object_group_id].page_bytes + alignment - 1) // alignment) + * alignment for key in lease.keys if key not in readable ) used, total = self._storage.get_l1_usage() + if reservation_failed and lease.unread_bytes > total: + return self._miss( + lease_id, lease, "cannot fit in RAM; its checkpoint stays listed" + ) # Without L2, or with more pages than RAM can hold, no repeat loads them. if ( self._storage.l2_adapters() and lease.unread_bytes <= total - and (lease.lookups == 1 or total - used < lease.unread_bytes) + and ( + reservation_failed + or lease.lookups == 1 + or total - used < lease.unread_bytes + ) ): lease.handle = None return self._repeat_with_room(lease_id, lease) @@ -692,16 +709,16 @@ def _repeat_with_room(self, lease_id: str, lease: _RetrieveLease) -> bool | None """ if lease.cancelled: return self._miss(lease_id, lease, "was cancelled by the engine") - used, total = self._storage.get_l1_usage() - if total - used >= lease.unread_bytes: - lease.handle = self._lookup(lease_id, lease.manifest, lease.keys) - lease.lookups += 1 - return None now = time.monotonic() if now - lease.started >= self._storage.store_admission_timeout_seconds: return self._miss( lease_id, lease, "found no room in RAM; its checkpoint stays listed" ) + used, total = self._storage.get_l1_usage() + if total - used >= lease.unread_bytes: + lease.handle = self._lookup(lease_id, lease.manifest, lease.keys) + lease.lookups += 1 + return None if now - lease.eviction_requested >= _ROOM_REQUEST_INTERVAL_SECONDS: lease.eviction_requested = now self._storage.request_immediate_eviction() diff --git a/tests/v1/multiprocess/test_checkpoint_restore_capacity.py b/tests/v1/multiprocess/test_checkpoint_restore_capacity.py new file mode 100644 index 00000000000..d57c163266d --- /dev/null +++ b/tests/v1/multiprocess/test_checkpoint_restore_capacity.py @@ -0,0 +1,165 @@ +# SPDX-License-Identifier: Apache-2.0 +"""RAM allocation failures must not delist intact disk checkpoints.""" + +# Standard +from dataclasses import replace +from pathlib import Path +from typing import Any +from unittest.mock import patch +import hashlib + +# Third Party +import pytest +import torch + +# First Party +from lmcache.v1.distributed.api import ( + MemoryLayoutDesc, + ObjectKey, + PrefetchHandle, + PrefetchRequestSpec, + TrimPolicy, +) +from lmcache.v1.distributed.storage_controllers.prefetch_controller import ( + PrefetchResult, +) +from lmcache.v1.multiprocess.checkpoint_storage import ( + checkpoint_object_keys, + checkpoint_page_groups, +) +from tests.v1.multiprocess.test_checkpoint_storage import ( + make_manifest, + open_store, + poll, + publish_before_restart, + restores, +) + + +@pytest.mark.parametrize("native", [False, True]) +def test_restore_accounts_for_page_alignment(tmp_path: Path, native: bool) -> None: + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native, admission_timeout_seconds=0.15) as ( + service, + index, + storage, + mapping, + ): + filler = ObjectKey(hashlib.sha256(b"alignment-filler").digest(), "filler", 0, 0) + _, total = storage.get_l1_usage() + alignment = storage.l1_memory_desc.align_bytes + reserved = storage.reserve_write_detailed( + [filler], + MemoryLayoutDesc([torch.Size([total - alignment])], [torch.uint8]), + "new", + ) + assert reserved[filler][1] is not None + try: + # Raw payload bytes fit; seven separate aligned allocations do not. + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + assert poll(service, lease_id) is False + assert index.get(entry.generation) == entry + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize("native", [False, True]) +def test_restore_uses_allocation_result_before_capacity_returns( + tmp_path: Path, native: bool +) -> None: + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native, admission_timeout_seconds=0.15) as ( + service, + index, + storage, + mapping, + ): + filler = ObjectKey( + hashlib.sha256(b"concurrent-writer").digest(), "filler", 0, 0 + ) + _, total = storage.get_l1_usage() + submit = storage.submit_prefetch_task + query = storage.query_prefetch_status_detailed + attempts = 0 + + def reserve_then_submit(*args: Any, **kwargs: Any) -> PrefetchHandle: + nonlocal attempts + attempts += 1 + result = storage.reserve_write_detailed( + [filler], MemoryLayoutDesc([torch.Size([total])], [torch.uint8]), "new" + ) + assert result[filler][1] is not None + return submit(*args, **kwargs) + + def release_before_consumption(handle: PrefetchHandle) -> PrefetchResult | None: + result = query(handle) + if result is not None: + # Enforce a concurrent writer finishing after allocation failed. + storage.abort_write([filler]) + return result + + try: + with ( + patch.object( + storage, "submit_prefetch_task", side_effect=reserve_then_submit + ), + patch.object( + storage, + "query_prefetch_status_detailed", + side_effect=release_before_consumption, + ), + ): + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + assert poll(service, lease_id) is False + assert attempts >= 2 + assert index.get(entry.generation) == entry + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize("native", [False, True]) +@pytest.mark.parametrize("detailed", [False, True]) +def test_prefetch_queries_consume_capacity_failure_once( + tmp_path: Path, native: bool, detailed: bool +) -> None: + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native) as (service, index, storage, mapping): + filler = ObjectKey(hashlib.sha256(b"query-filler").digest(), "filler", 0, 0) + _, total = storage.get_l1_usage() + reserved = storage.reserve_write_detailed( + [filler], MemoryLayoutDesc([torch.Size([total])], [torch.uint8]), "new" + ) + assert reserved[filler][1] is not None + keys = [key for group in checkpoint_object_keys(entry, 0) for key in group] + layouts = { + i: MemoryLayoutDesc([torch.Size([group.page_bytes])], [torch.uint8]) + for i, group in enumerate(checkpoint_page_groups(entry)) + } + try: + handle = storage.submit_prefetch_task( + PrefetchRequestSpec(keys, layouts, policy=TrimPolicy.SPARSE) + ) + assert storage.wait_prefetch_status(handle, timeout=5) + if detailed: + result = storage.query_prefetch_status_detailed(handle) + assert result is not None and result.reservation_failed + found = result.found + else: + legacy_found = storage.query_prefetch_status(handle) + assert legacy_found is not None + found = legacy_found + assert found is not None and found.popcount() == 0 + assert storage.query_prefetch_status_detailed(handle) is None + assert storage.query_prefetch_status(handle) is None + assert storage.query_prefetch_lookup_hits(handle) is None + assert index.get(entry.generation) == entry + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry)