From 6d6f8e6afcd1be4f82735d8a5226438f7e086ee6 Mon Sep 17 00:00:00 2001 From: Costa Tsaousis Date: Tue, 29 Sep 2026 09:50:26 +0000 Subject: [PATCH 1/3] fix(checkpoint): preserve disk checkpoints after allocation failure --- .../prefetch_controller.py | 43 ++++- lmcache/v1/distributed/storage_manager.py | 23 ++- lmcache/v1/multiprocess/checkpoint_storage.py | 45 +++-- .../test_checkpoint_restore_capacity.py | 159 ++++++++++++++++++ 4 files changed, 248 insertions(+), 22 deletions(-) create mode 100644 tests/v1/multiprocess/test_checkpoint_restore_capacity.py 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..52cbd5744f1 --- /dev/null +++ b/tests/v1/multiprocess/test_checkpoint_restore_capacity.py @@ -0,0 +1,159 @@ +# 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. + assert poll(service, service.begin_retrieve(entry, 0)) 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, + ), + ): + assert poll(service, service.begin_retrieve(entry, 0)) 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: + found = storage.query_prefetch_status(handle) + 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) From 8a2d97b4ce9414622ae23082336abccf8d5029cc Mon Sep 17 00:00:00 2001 From: Costa Tsaousis Date: Tue, 29 Sep 2026 10:00:27 +0000 Subject: [PATCH 2/3] test(checkpoint): narrow optional restore results before use --- .../multiprocess/test_checkpoint_restore_capacity.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/tests/v1/multiprocess/test_checkpoint_restore_capacity.py b/tests/v1/multiprocess/test_checkpoint_restore_capacity.py index 52cbd5744f1..d57c163266d 100644 --- a/tests/v1/multiprocess/test_checkpoint_restore_capacity.py +++ b/tests/v1/multiprocess/test_checkpoint_restore_capacity.py @@ -57,7 +57,9 @@ def test_restore_accounts_for_page_alignment(tmp_path: Path, native: bool) -> No assert reserved[filler][1] is not None try: # Raw payload bytes fit; seven separate aligned allocations do not. - assert poll(service, service.begin_retrieve(entry, 0)) is False + 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]) @@ -111,7 +113,9 @@ def release_before_consumption(handle: PrefetchHandle) -> PrefetchResult | None: side_effect=release_before_consumption, ), ): - assert poll(service, service.begin_retrieve(entry, 0)) is False + 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: @@ -148,7 +152,9 @@ def test_prefetch_queries_consume_capacity_failure_once( assert result is not None and result.reservation_failed found = result.found else: - found = storage.query_prefetch_status(handle) + 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 From 0337b0e33368c8cdd1547bcedf7eb3413bcc1c58 Mon Sep 17 00:00:00 2001 From: Martin Vit Date: Tue, 29 Sep 2026 19:36:49 +0000 Subject: [PATCH 3/3] fix(checkpoint): bound #105's RAM-failure repeats and keep retiring lost pages On top of ktsaou's #105: - the admission timeout bounds only repeats after a RAM reservation failure; other misses repeat once when there is room, as before, so a checkpoint with a deleted page file is still retired when the first lookup answers after the timeout or the timeout is zero; - a repeat after a RAM failure waits 20 ms instead of resubmitting on every poll (700-900 disk lookups per second before); - #106's stalled-lookup test helper patches the detailed prefetch query the store now uses; - regression tests for a writer that frees RAM before the poll, fragmented RAM (listing kept, resubmits bounded) and lost pages under slow or zero timeouts; release fragment lmcache-105. Co-Authored-By: Claude Opus 5.5 --- .lil/changes/lmcache-105.json | 23 +++ lmcache/v1/multiprocess/checkpoint_storage.py | 20 +- .../test_checkpoint_cancel_drain.py | 4 +- .../test_checkpoint_restore_ram_failures.py | 180 ++++++++++++++++++ 4 files changed, 220 insertions(+), 7 deletions(-) create mode 100644 .lil/changes/lmcache-105.json create mode 100644 tests/v1/multiprocess/test_checkpoint_restore_ram_failures.py diff --git a/.lil/changes/lmcache-105.json b/.lil/changes/lmcache-105.json new file mode 100644 index 00000000000..86979be1488 --- /dev/null +++ b/.lil/changes/lmcache-105.json @@ -0,0 +1,23 @@ +{ + "schema": "local-inference-release-change/v1", + "id": "lmcache-105", + "category": "fix", + "summary": "A checkpoint restore that cannot get RAM no longer removes an intact disk checkpoint from the directory", + "models": [ + "GLM-5.3-Flash", + "Qwen3.8-Flash-Next" + ], + "compatibility": "No user action required.", + "details": [ + "A restore from disk loads a checkpoint's pages only into RAM reserved for all of them. When that reservation failed, the restore was repeated once, and if the repeat failed too while RAM looked free by the time the cache checked (RAM held briefly by other restores or writes, or free RAM split into pieces smaller than a page), the checkpoint was removed from the directory although every page was intact on disk. Later turns of that conversation then recomputed their prompts instead of restoring.", + "The cache now decides from the allocation result itself: a restore whose pages could not be reserved is repeated, 20 ms apart, until the store admission timeout (default 8 s), and then misses with its checkpoint still listed, so a later request restores it. Only a restore that got its RAM and still found pages unreadable retires the checkpoint, also when the first lookup answers after the timeout. A checkpoint larger than the whole RAM tier also stays listed, and the room a waiting restore asks eviction for now includes each page's 4 KiB alignment." + ], + "pull_requests": [ + 105 + ], + "requires": [], + "authors": [ + "ktsaou", + "voipmonitor" + ] +} diff --git a/lmcache/v1/multiprocess/checkpoint_storage.py b/lmcache/v1/multiprocess/checkpoint_storage.py index 3c5f117de40..9c8ef9477e2 100644 --- a/lmcache/v1/multiprocess/checkpoint_storage.py +++ b/lmcache/v1/multiprocess/checkpoint_storage.py @@ -49,6 +49,8 @@ # seconds, as store admission does: a pass frees nothing while its victims are # still pinned by other restores or L2 writes. _ROOM_REQUEST_INTERVAL_SECONDS = 0.5 +# Pause before repeating a lookup whose load buffers could not be reserved. +_RESERVATION_RETRY_SECONDS = 0.02 # Released cancelled lookups remembered so that a later poll from their worker # still gets the miss; a poll of an older one raises KeyError. _MAX_RELEASED_CANCELLED_LOOKUPS = 4096 @@ -237,6 +239,7 @@ class _RetrieveLease: readable: int = 0 unread_bytes: int = 0 eviction_requested: float = 0.0 + resubmit_after: float = 0.0 class CheckpointPayloadStore: @@ -743,6 +746,13 @@ def _repeat_or_miss( return self._miss( lease_id, lease, "cannot fit in RAM; its checkpoint stays listed" ) + if reservation_failed: + 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" + ) + lease.resubmit_after = now + _RESERVATION_RETRY_SECONDS # Without L2, or with more pages than RAM can hold, no repeat loads them. if ( self._storage.l2_adapters() @@ -769,15 +779,15 @@ 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") 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: + if total - used >= lease.unread_bytes and now >= lease.resubmit_after: lease.handle = self._lookup(lease_id, lease.manifest, lease.keys) lease.lookups += 1 return None + 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" + ) 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_cancel_drain.py b/tests/v1/multiprocess/test_checkpoint_cancel_drain.py index df8e4e305c9..5c0c5aa3fb5 100644 --- a/tests/v1/multiprocess/test_checkpoint_cancel_drain.py +++ b/tests/v1/multiprocess/test_checkpoint_cancel_drain.py @@ -141,10 +141,10 @@ def test_failed_poll_of_a_pending_lookup_is_not_fatal() -> None: def _stall_lookups(monkeypatch: pytest.MonkeyPatch, storage: Any) -> threading.Event: """Keep every prefetch pending until the returned event is set.""" answer = threading.Event() - query = storage.query_prefetch_status + query = storage.query_prefetch_status_detailed monkeypatch.setattr( storage, - "query_prefetch_status", + "query_prefetch_status_detailed", lambda handle: query(handle) if answer.is_set() else None, ) return answer diff --git a/tests/v1/multiprocess/test_checkpoint_restore_ram_failures.py b/tests/v1/multiprocess/test_checkpoint_restore_ram_failures.py new file mode 100644 index 00000000000..deb76b020f5 --- /dev/null +++ b/tests/v1/multiprocess/test_checkpoint_restore_ram_failures.py @@ -0,0 +1,180 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Checkpoint restores whose RAM reservation fails keep intact checkpoints listed. + +A restore from disk reserves RAM for all of a checkpoint's pages. When that +reservation fails, the checkpoint must stay listed so a later request restores +it; only pages that are really missing retire it. +""" + +# Standard +from dataclasses import replace +from pathlib import Path +from typing import Any +from unittest.mock import patch +import hashlib +import time + +# Third Party +import pytest +import torch + +# First Party +from lmcache.v1.distributed.api import MemoryLayoutDesc, ObjectKey +from tests.v1.multiprocess.test_checkpoint_storage import ( + entry_keys, + large_manifest, + make_manifest, + open_store, + poll, + publish_before_restart, + restores, +) + + +@pytest.mark.parametrize("native", [False, True]) +def test_ram_released_before_the_poll_keeps_checkpoint_listed( + tmp_path: Path, native: bool +) -> None: + """A writer holds all RAM while the load buffers are reserved and frees it + before the store reads the result: the checkpoint stays listed.""" + 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"probe-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) -> Any: + 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: Any) -> Any: + result = query(handle) + if result is not None: + 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 + outcome = poll(service, lease_id) + listed = index.get(entry.generation) == entry + assert outcome is False + assert listed, "intact disk checkpoint was delisted" + finally: + storage.abort_write([filler]) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize("native", [False, True]) +def test_fragmented_ram_keeps_checkpoint_listed_without_spinning( + tmp_path: Path, native: bool +) -> None: + """Free RAM exceeds the unread pages, but no free block holds one page. + + No concurrency: 64 KiB holes between pinned 64 KiB blocks, 128 KiB + pages. The reservation fails every time while get_l1_usage shows room, so + the lookup is repeated with a pause until the admission timeout. + """ + entry = large_manifest(500) + publish_before_restart(tmp_path, native, [entry]) + with open_store(tmp_path, native, admission_timeout_seconds=1.0) as ( + service, + index, + storage, + mapping, + ): + used, total = storage.get_l1_usage() + assert used == 0 + block = 64 * 1024 + fillers = [ + ObjectKey(hashlib.sha256(b"frag%d" % i).digest(), "filler", 0, 0) + for i in range(total // block) + ] + reserved = storage.reserve_write_detailed( + fillers, MemoryLayoutDesc([torch.Size([block])], [torch.uint8]), "new" + ) + assert all(obj is not None for _err, obj in reserved.values()) + storage.abort_write(fillers[::2]) + kept = fillers[1::2] + used, total = storage.get_l1_usage() + unread = len(entry_keys(entry)) * 128 * 1024 + assert total - used >= unread + submit = storage.submit_prefetch_task + attempts = 0 + + def count(*args: Any, **kwargs: Any) -> Any: + nonlocal attempts + attempts += 1 + return submit(*args, **kwargs) + + try: + with patch.object(storage, "submit_prefetch_task", side_effect=count): + start = time.monotonic() + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + outcome = poll(service, lease_id) + elapsed = time.monotonic() - start + listed = index.get(entry.generation) == entry + assert outcome is False + assert listed, "intact disk checkpoint was delisted" + # Repeats after a RAM failure pause 20 ms: about 50 in one second. + assert attempts <= 80, f"{attempts} lookups in {elapsed:.2f} s" + finally: + storage.abort_write(kept) + assert restores(service, mapping, entry) + + +@pytest.mark.parametrize( + "timeout,delay", + [(0.2, 0.0), (0.2, 0.4), (0.0, 0.0)], + ids=["fast-first-lookup", "first-lookup-slower-than-timeout", "timeout-0"], +) +def test_lost_page_retires_checkpoint( + tmp_path: Path, timeout: float, delay: float +) -> None: + """A page file deleted from disk (not a RAM problem) retires its checkpoint. + + Also when the first lookup is answered after the admission timeout, or the + timeout is zero: the timeout bounds only repeats after a RAM failure. + """ + entry = replace(make_manifest(), world_size=1) + publish_before_restart(tmp_path, False, [entry]) + files = sorted((tmp_path / "payloads").rglob("*.data")) + assert len(files) == len(entry_keys(entry)) + files[0].unlink() + with open_store(tmp_path, False, admission_timeout_seconds=timeout) as ( + service, + index, + storage, + _mapping, + ): + lease_id = service.begin_retrieve(entry, 0) + assert lease_id is not None + time.sleep(delay) + outcome = poll(service, lease_id) + listed = index.get(entry.generation) == entry + assert outcome is False + assert not listed, "checkpoint with a lost page stays listed"