diff --git a/.lil/changes/lmcache-104.json b/.lil/changes/lmcache-104.json new file mode 100644 index 00000000000..688a0661c8b --- /dev/null +++ b/.lil/changes/lmcache-104.json @@ -0,0 +1,25 @@ +{ + "schema": "local-inference-release-change/v1", + "id": "lmcache-104", + "category": "fix", + "summary": "Conversations restore after a title request or sub-agent fork, lost checkpoints are no longer offered, and a stopping container finishes and writes its last checkpoints", + "models": [ + "Qwen3.8-Flash-Next", + "GLM-5.3-Flash" + ], + "compatibility": "No user action required. New option --checkpoint-supersede-grace-seconds (default 300). --checkpoint-shutdown-flush-seconds now also covers a short wait (at most a third of it and 5 s) for checkpoint stores still in flight when the container stops. Give the container time to stop (docker run --stop-timeout 60 or Compose stop_grace_period: 60s); with Docker's default 10 s the current checkpoints are written first.", + "details": [ + "A request that extends a conversation's latest checkpoint without being its next turn, such as a title, summary or follow-up request, or a sub-agent forked from the conversation, marked that checkpoint as superseded. Under RAM pressure its state pages were dropped without a disk write, and when the user continued, the restore failed with 'K of M pages were readable' and recomputed the prompt. The checkpoint a new prompt continues from now stays current until a later prompt moves past it too, so the next turn restores it. Superseded checkpoints are still never written to disk while serving.", + "A fork point that a sub-agent superseded keeps its place in RAM eviction order for 300 s instead of being dropped first, and a lookup that finds it makes it current again.", + "A checkpoint is retired from the directory before the last copy of a page it needs is dropped, and a lookup skips and retires checkpoints whose pages no tier holds (evicted from disk, deleted, or not written before a restart). The next shorter checkpoint is restored in the same lookup instead of after a failed restore. Pages that another server sharing the cache directory wrote still count.", + "When the container stops, the cache keeps serving the checkpoint stores the model is still copying for a few seconds, then writes checkpoints that are only in RAM, current ones first, and logs what it could not write. Before, a store in flight at that moment made the cache skip the whole shutdown write.", + "In a synthetic 12-conversation workload with side requests, failed restores fell from 118-247 per 180 turns to 0. Disk writes grow only where RAM barely holds the active conversations, by at most about one checkpoint per turn." + ], + "pull_requests": [ + 104 + ], + "evidence": [ + "https://github.com/mrweiner/lmcache-supersession-repro" + ], + "requires": [] +} diff --git a/docs/design/v1/distributed/l2_adapters/overall.md b/docs/design/v1/distributed/l2_adapters/overall.md index 2ce21b8a7e3..d5dbd708d7d 100644 --- a/docs/design/v1/distributed/l2_adapters/overall.md +++ b/docs/design/v1/distributed/l2_adapters/overall.md @@ -528,6 +528,15 @@ existing `close()` (to keep data on disk) path. See `nixl_store_dynamic_l2_adapter.py` for a reference implementation and [`nixl_store.md`](nixl_store.md) for design details. +`absent_keys(keys)` returns the keys an adapter can confirm it does not +hold, answered synchronously from cheap local state; the default confirms +nothing. The storage manager treats a recurrent checkpoint page as lost only +when L1 does not hold it and every adapter confirms it absent, and a +checkpoint lookup then retires the checkpoints that need it instead of +offering a restore that fails. The check must see objects that other +processes sharing the backend wrote: `fs` and `fs_native` stat the object +files, and the in-process mock checks its dictionary. + ### Native (C++/Rust) Storage Backends For high-performance backends written in C++ or Rust, use the shared native diff --git a/docs/design/v1/mp_observability/METRICS.md b/docs/design/v1/mp_observability/METRICS.md index 6dbc03cfece..37a38269c96 100644 --- a/docs/design/v1/mp_observability/METRICS.md +++ b/docs/design/v1/mp_observability/METRICS.md @@ -471,7 +471,7 @@ the same backend type — same shape as the existing |---|---|---|---|---| | `lmcache_mp.l1_memory_usage_bytes` | `lmcache_mp_l1_memory_usage_bytes` | ObservableGauge | `L1Manager.get_memory_usage()` | Bytes currently held in L1 at scrape time | | `lmcache_mp.l2_usage_bytes` | `lmcache_mp_l2_usage_bytes` | ObservableGauge (attr: `l2_name`) | `StorageManager.get_l2_usages()` (calls `L2AdapterInterface.get_usage().total_bytes_used`) | Per-adapter bytes currently held in L2 at scrape time; one observation per configured adapter. Adapters whose `get_usage()` raises are skipped silently. | -| `lmcache_mp.checkpoint_retention` | `lmcache_mp_checkpoint_retention` | ObservableGauge (attr: `stat`) | `CheckpointRetention.observations()` | Recurrent checkpoint retention: `superseded_checkpoints`, `superseded_pages`, `superseded_pages_tracked`, `l1_superseded_drops`, `l2_superseded_evictions`, `l2_superseded_eviction_bytes`, `write_on_evict_requests`, `write_on_evict_persisted`, `write_on_evict_timeouts`, `write_on_evict_pending`, `l2_resident_pages_tracked`, `l2_checkpoint_bytes`. Counters are cumulative since start. | +| `lmcache_mp.checkpoint_retention` | `lmcache_mp_checkpoint_retention` | ObservableGauge (attr: `stat`) | `CheckpointRetention.observations()` | Recurrent checkpoint retention: `superseded_checkpoints`, `superseded_pages`, `superseded_pages_tracked`, `superseded_pages_in_grace`, `l1_superseded_drops`, `l2_superseded_evictions`, `l2_superseded_eviction_bytes`, `retired_checkpoints`, `restored_checkpoints`, `write_on_evict_requests`, `write_on_evict_persisted`, `write_on_evict_timeouts`, `write_on_evict_pending`, `l2_resident_pages_tracked`, `l2_checkpoint_bytes`. Counters are cumulative since start. | | `lmcache_mp.num_inflight_l2_stores` | `lmcache_mp_num_inflight_l2_stores` | ObservableGauge (attrs: `l2_name`, `adapter_index`) | `StoreController.get_inflight_count_by_adapter()` | Snapshot of in-flight L2 store tasks grouped by adapter | | `lmcache_mp.num_inflight_l2_loads` | `lmcache_mp_num_inflight_l2_loads` | ObservableGauge (attrs: `l2_name`, `adapter_index`) | `PrefetchController.get_inflight_load_state_by_adapter()` | Per-adapter count from the same snapshot | | `lmcache_mp.inflight_load_memory_usage_bytes` | `lmcache_mp_inflight_load_memory_usage_bytes` | ObservableGauge (attrs: `l2_name`, `adapter_index`) | `PrefetchController.get_inflight_load_state_by_adapter()` | Per-adapter reserved bytes from the same snapshot | diff --git a/docs/source/mp/configuration.rst b/docs/source/mp/configuration.rst index 9a41433523a..d54f2857e24 100644 --- a/docs/source/mp/configuration.rst +++ b/docs/source/mp/configuration.rst @@ -493,10 +493,21 @@ Source: ``lmcache/v1/distributed/config.py`` evicted without an L2 copy and a warning is logged. * - ``--checkpoint-shutdown-flush-seconds`` - ``30`` - - With ``checkpoint_on_evict``, how long shutdown waits while current - checkpoint pages still only in L1 are written to L2, so the next - start can restore them. ``0`` skips the flush. Give the process - at least this much time between SIGTERM and SIGKILL. + - Budget of a clean shutdown for recurrent checkpoints. The server + first keeps serving checkpoint stores still in flight (at most a + third of the budget, at most 5 s, and only while stores are in + flight), so an engine stopping at the same time can publish them. + With ``checkpoint_on_evict`` it then writes checkpoint pages still + only in L1 to L2, current pages first, so the next start can restore + them. ``0`` skips both. Give the process at least this much time + between SIGTERM and SIGKILL. + * - ``--checkpoint-supersede-grace-seconds`` + - ``300`` + - How long a superseded recurrent checkpoint that a later prompt had + continued from (a branch point, such as a turn a sub-agent forked + from) keeps its eviction order before its pages are dropped first. + Superseded pages are never written to L2 while serving either way. + ``0`` drops them first at once. * - ``--l2-prefetch-policy`` - ``default`` - L2 prefetch policy. Determines which adapter loads each key diff --git a/docs/source/mp/l2_storage/index.rst b/docs/source/mp/l2_storage/index.rst index a89c8c3a188..cfdac3e8eab 100644 --- a/docs/source/mp/l2_storage/index.rst +++ b/docs/source/mp/l2_storage/index.rst @@ -450,29 +450,64 @@ auxiliary state are unique to each. The vLLM integration reports each published checkpoint with its token sequence and request. When the prompt checkpoint of a new request is -published, the server marks as *superseded* the unique pages of every -shorter checkpoint of that sequence produced by an earlier request, and of -the other checkpoints those requests published. That includes a response -endpoint that the chat template rewrote and the next prompt therefore does -not extend. Checkpoints of the same request and ``instruction`` -checkpoints, which other conversations share, are never marked. With every -store policy: - -* L1 eviction drops superseded pages before any LRU victim. -* L2 eviction deletes superseded pages before any LRU victim. +published, the longest earlier checkpoint that it extends is where it +continues from. That checkpoint stays current: the new request may be a +side request (a title or summary request that appends a task to the whole +conversation, or a sub-agent forked from it) while the conversation itself +goes on from the same point later. The server marks as *superseded* the +unique pages of every other shorter checkpoint of that sequence produced by +an earlier request (the conversation has moved past them twice), and of the +other checkpoints those requests published that the new prompt is longer +than. That includes a response endpoint that the chat template rewrote and +the next prompt therefore does not extend. A checkpoint at least as long as +the new prompt is kept: the new request branched off before it. +Checkpoints of the same request and ``instruction`` checkpoints, which other +conversations share, are never marked. With every store policy: + +* L1 eviction drops stale superseded pages before any LRU victim. +* L2 eviction deletes stale superseded pages before any LRU victim. + +A superseded page is stale at once, unless a later prompt had continued from +its checkpoint (a possible branch point, for example the end of a turn that +a sub-agent forked from and ran several turns on). Such pages keep their +normal eviction order for ``--checkpoint-supersede-grace-seconds`` (default +300), so the conversation can still continue from them. A lookup that finds +a superseded checkpoint makes it current again. With ``--l2-store-policy checkpoint_on_evict`` in addition: * A current checkpoint page is written to L2 once, when L1 is about to evict it, instead of on every request. L1 keeps the page until the write completes (at most ``--checkpoint-write-timeout-seconds``). -* A superseded page is never written to L2. -* Shutdown writes current pages still only in L1, so they can be restored - after a restart. - -Superseded manifests stay listed, so a request that branches from an older -turn can still restore it while its pages last; when a page is gone, the -restore falls back to the longest remaining checkpoint. +* A superseded page is never written to L2 while serving, also during its + grace period. +* Shutdown writes pages still only in L1, current pages first, so they can + be restored after a restart. + +A checkpoint is never offered once a page it needs is gone. Before the last +copy of a superseded page is deleted from L1 or L2, the superseded +checkpoints that reference it are retired from the directory. A lookup also +checks the pages of the longest candidate against L1 and the L2 inventory: +when a page was lost another way (L2 eviction, an administrative delete, a +page that was not written before a restart), the lookup retires that +candidate and returns the next shorter complete one in the same call. A page +counts as lost only when every L2 adapter confirms it does not hold it; +``fs`` and ``fs_native`` check the object file, so pages another server +sharing the directory wrote still count. With adapters that cannot confirm +an absence, the restore still validates every page and falls back to a +shorter checkpoint. + +**Shutdown.** ``--checkpoint-shutdown-flush-seconds`` (default 30) is the +budget of a clean shutdown. At SIGTERM the server first keeps serving +checkpoint stores that are still in flight, so an engine stopping at the +same time can publish the checkpoints it is copying. This takes at most a +third of the budget (and at most 5 s) and ends as soon as no store is in +flight. With ``checkpoint_on_evict`` it then writes pages still only in L1 +to L2 within the rest of the budget, current pages first, and logs how many +it could not write. Give the process that much time between SIGTERM and +SIGKILL; with a shorter stop timeout the most valuable pages are written +first, and the checkpoints left incomplete are retired by lookups after the +restart. **Sizing.** With write-through (``default``) the L2 retention time is about:: @@ -485,5 +520,7 @@ per conversation that leaves L1, so retention is about:: and superseded pages are reclaimed first. The ``lmcache_mp_checkpoint_retention`` gauge reports, per ``stat``, superseded -pages, L1 drops and L2 evictions of superseded pages, write-on-evict -requests, completions and timeouts, and the checkpoint bytes held in L2. +pages (and those in their grace period), L1 drops and L2 evictions of +superseded pages, retired checkpoints, superseded checkpoints found again, +write-on-evict requests, completions and timeouts, and the checkpoint bytes +held in L2. diff --git a/lmcache/v1/distributed/checkpoint_retention.py b/lmcache/v1/distributed/checkpoint_retention.py index 1cb29a77af3..0f03afd1e71 100644 --- a/lmcache/v1/distributed/checkpoint_retention.py +++ b/lmcache/v1/distributed/checkpoint_retention.py @@ -8,22 +8,33 @@ checkpoint (recurrent endpoint state, a partial attention page, auxiliary state) become dead weight. This module keeps that knowledge: -* Superseded pages are dropped first from L1 and L2 and are never written to - L2 when they leave L1. +* Superseded pages are never written to L2 when they leave L1. Once stale + they are dropped first from L1 and L2. A page is stale as soon as it is + superseded, unless its checkpoint is one a later prompt continued from: a + branch point such as the end of a turn that a sub-agent forked from and + ran several turns on. For a grace period such pages keep their LRU + position instead, so the original line can still continue from them + while they are in RAM. +* Before the last copy of a superseded page is deleted, the superseded + checkpoints that need it are retired from the directory, so a lookup + misses them cleanly and falls back to a shorter checkpoint at once. * With write-on-evict storage, a current checkpoint page is written to L2 once, when L1 is about to evict it, instead of on every request. * L2 residency per adapter is tracked so a page that already reached L2 is - not written again. + not written again. A lookup uses it to find pages that may be lost, and + treats one as lost only when L1 does not hold it and every L2 adapter + confirms that it does not either. Everything is bounded in memory and safe to call from the store, eviction and -request-handler threads. Forgetting an entry only costs an extra write or a -later eviction; correctness never depends on this state, because restores +request-handler threads. Forgetting a supersession only costs an extra write +or a later eviction. Where no adapter can confirm an absence, restores still validate every page and fall back to a shorter checkpoint when one is gone. """ # Standard from collections import OrderedDict from collections.abc import Callable, Iterable +from dataclasses import dataclass, field import threading import time @@ -57,8 +68,30 @@ def __contains__(self, item: object) -> bool: def __len__(self) -> int: return len(self._items) - def items(self) -> list: - return list(self._items) + +@dataclass +class _SupersededPage: + """A superseded page: when it may be dropped first, and who needs it.""" + + stale_at: float + """Monotonic time from which the page is dropped first.""" + + generations: set[str] = field(default_factory=set) + """Superseded checkpoints that reference the page.""" + + +def _all_present(keys: list[ObjectKey]) -> list[ObjectKey]: + """Default L1 lookup: without one, assume L1 may hold every page.""" + return keys + + +def _retire_nothing(generations: list[str]) -> None: + """Default retirement hook: without a directory there is nothing to delist.""" + + +def _all_absent(keys: list[ObjectKey]) -> list[ObjectKey]: + """Default L2 absence check: without L2 adapters no page is in L2.""" + return keys class _AdapterResidencyListener(L2AdapterListener): @@ -90,8 +123,13 @@ class CheckpointRetention: the storage manager once its store controller exists. max_superseded: Bound on remembered superseded pages. max_tracked: Bound on remembered L2-resident pages per adapter. + max_generations: Bound on remembered superseded and continued + checkpoints. persist_timeout: Seconds after which a page whose write never completed may be evicted from L1 without reaching L2. + supersede_grace_seconds: How long a superseded checkpoint that a + later prompt continued from keeps its LRU position before its + pages are dropped first. 0 drops every superseded page first. """ def __init__( @@ -103,18 +141,33 @@ def __init__( max_tracked: int = 1048576, max_generations: int = 65536, persist_timeout: float = 120.0, + supersede_grace_seconds: float = 0.0, ) -> None: + if supersede_grace_seconds < 0: + raise ValueError("supersede_grace_seconds must be >= 0") self._write_on_evict = write_on_evict self._persist = persist + self._max_superseded = max_superseded self._max_tracked = max_tracked + self._max_generations = max_generations self._persist_timeout = persist_timeout + self._grace = supersede_grace_seconds self._lock = threading.Lock() - self._superseded = _BoundedSet(max_superseded) - self._superseded_generations = _BoundedSet(max_generations) + # Superseded pages in the order they were first marked. + self._superseded: OrderedDict[ObjectKey, _SupersededPage] = OrderedDict() + # Superseded checkpoints and the pages they alone referenced. + self._superseded_generations: OrderedDict[str, tuple[ObjectKey, ...]] = ( + OrderedDict() + ) + # Checkpoints a later prompt continued from. + self._continued = _BoundedSet(max_generations) # adapter id -> key -> size in bytes self._resident: dict[int, OrderedDict[ObjectKey, int]] = {} # key -> monotonic time its write was requested self._pending: dict[ObjectKey, float] = {} + self._l1_present: Callable[[list[ObjectKey]], list[ObjectKey]] = _all_present + self._l2_absent: Callable[[list[ObjectKey]], list[ObjectKey]] = _all_absent + self._retire: Callable[[list[str]], None] = _retire_nothing self._stats = { "superseded_checkpoints": 0, "superseded_pages": 0, @@ -124,6 +177,8 @@ def __init__( "write_on_evict_requests": 0, "write_on_evict_persisted": 0, "write_on_evict_timeouts": 0, + "retired_checkpoints": 0, + "restored_checkpoints": 0, } @property @@ -134,6 +189,46 @@ def set_persist(self, persist: Callable[[list[ObjectKey]], None]) -> None: """Install the asynchronous L2 store callback.""" self._persist = persist + def set_l1_lookup( + self, present: Callable[[list[ObjectKey]], list[ObjectKey]] + ) -> None: + """Install the L1 lookup used to decide whether a copy is the last. + + Args: + present: Returns the given keys that L1 holds. It must not call + back into this object. + """ + self._l1_present = present + + def set_l2_absence( + self, absent: Callable[[list[ObjectKey]], list[ObjectKey]] + ) -> None: + """Install the check that confirms pages are in no L2 adapter. + + Args: + absent: Returns the given keys that every L2 adapter confirms it + does not hold (all of them without adapters). It must not + call back into this object. + """ + self._l2_absent = absent + + def set_retirement(self, retire: Callable[[list[str]], None]) -> None: + """Install the hook that delists checkpoints whose pages are lost. + + Args: + retire: Removes the given generations from the checkpoint + directory. It is called without this object's lock held, + before the page deletion that loses them, and must not raise. + """ + with self._lock: + self._retire = retire + + def clear_retirement(self, retire: Callable[[list[str]], None]) -> None: + """Remove ``retire`` if it is still the installed retirement hook.""" + with self._lock: + if self._retire == retire: + self._retire = _retire_nothing + def listener_for(self, adapter_id: int) -> L2AdapterListener: """Return a listener that records one adapter's checkpoint pages.""" with self._lock: @@ -146,6 +241,15 @@ def forget_adapter(self, adapter_id: int) -> None: # ----- supersession --------------------------------------------------- + def mark_continued(self, generation: str) -> None: + """Record that a later prompt directly continued from ``generation``. + + Such a checkpoint may be a branch point: if it is superseded later, + its pages keep their eviction order for the grace period. + """ + with self._lock: + self._continued.add(generation) + def is_superseded_generation(self, generation: str) -> bool: with self._lock: return generation in self._superseded_generations @@ -153,21 +257,37 @@ def is_superseded_generation(self, generation: str) -> bool: def mark_superseded(self, generation: str, keys: Iterable[ObjectKey]) -> int: """Mark one older checkpoint's unique pages as superseded. + The pages become stale at once, or after the grace period when a + later prompt continued from ``generation``. + Args: generation: The older checkpoint's generation, remembered so the - same ancestor is not processed again. + same checkpoint is not processed again and can be retired + when one of its pages is lost. keys: Pages that no current checkpoint references. Returns: Number of pages newly marked. """ + now = time.monotonic() added = 0 with self._lock: - self._superseded_generations.add(generation) - for key in keys: - if key not in self._superseded: + stale_at = now + self._grace if generation in self._continued else now + pages = tuple(keys) + self._superseded_generations[generation] = pages + self._superseded_generations.move_to_end(generation) + while len(self._superseded_generations) > self._max_generations: + self._superseded_generations.popitem(last=False) + for key in pages: + page = self._superseded.get(key) + if page is None: + self._superseded[key] = _SupersededPage(stale_at, {generation}) added += 1 - self._superseded.add(key) + else: + page.generations.add(generation) + page.stale_at = max(page.stale_at, stale_at) + while len(self._superseded) > self._max_superseded: + self._superseded.popitem(last=False) self._stats["superseded_checkpoints"] += 1 self._stats["superseded_pages"] += added return added @@ -176,15 +296,147 @@ def mark_current(self, keys: Iterable[ObjectKey]) -> None: """Clear the superseded mark from pages a new checkpoint references.""" with self._lock: for key in keys: - self._superseded.discard(key) + self._superseded.pop(key, None) + + def mark_generation_current( + self, generation: str, keys: Iterable[ObjectKey] + ) -> bool: + """Make a superseded checkpoint current again. + + Used when a lookup finds the checkpoint: a request continues from it, + so it is no longer dead weight, and a later prompt may supersede it + again. + + Args: + generation: The found checkpoint. + keys: All of its pages. + + Returns: + True if the checkpoint was superseded. + """ + with self._lock: + was_superseded = ( + self._superseded_generations.pop(generation, None) is not None + ) + for key in keys: + self._superseded.pop(key, None) + if was_superseded: + self._stats["restored_checkpoints"] += 1 + return was_superseded def is_superseded(self, key: ObjectKey) -> bool: with self._lock: return key in self._superseded def superseded_keys(self) -> list[ObjectKey]: + """Return every superseded page, stale or not, oldest mark first.""" with self._lock: - return self._superseded.items() + return list(self._superseded) + + def take_droppable_superseded(self) -> list[ObjectKey]: + """Return stale superseded pages that L1 holds, oldest mark first. + + Stale pages that no tier holds any more are forgotten. + """ + now = time.monotonic() + with self._lock: + stale = [ + key for key, page in self._superseded.items() if page.stale_at <= now + ] + if not stale: + return [] + present = set(self._l1_present(stale)) + with self._lock: + for key in stale: + if key not in present and not self._is_resident_locked(key): + self._superseded.pop(key, None) + return [key for key in stale if key in present] + + # ----- retirement ----------------------------------------------------- + + def retire_before_l1_delete(self, keys: list[ObjectKey]) -> int: + """Retire superseded checkpoints that deleting ``keys`` from L1 loses. + + Call it before L1 deletes ``keys``. A superseded page without an L2 + copy is its last copy, so every superseded checkpoint that references + it is delisted first. Ordinary KV chunks return at once. + + Args: + keys: Keys L1 is about to delete. + + Returns: + Number of checkpoints retired. + """ + tracked = self._tracked_superseded(keys) + if not tracked: + return 0 + with self._lock: + lost = [key for key in tracked if not self._is_resident_locked(key)] + return self._retire_owners(lost) + + def retire_before_l2_delete(self, adapter_id: int, keys: list[ObjectKey]) -> int: + """Retire superseded checkpoints that deleting ``keys`` from L2 loses. + + Call it before adapter ``adapter_id`` deletes ``keys``. A superseded + page that neither L1 nor another adapter holds is its last copy. + + Args: + adapter_id: The adapter about to delete ``keys``. + keys: Keys it is about to delete. + + Returns: + Number of checkpoints retired. + """ + tracked = self._tracked_superseded(keys) + if not tracked: + return 0 + in_l1 = set(self._l1_present(tracked)) + with self._lock: + lost = [ + key + for key in tracked + if key not in in_l1 + and not any( + key in resident + for other, resident in self._resident.items() + if other != adapter_id + ) + ] + return self._retire_owners(lost) + + def _tracked_superseded(self, keys: list[ObjectKey]) -> list[ObjectKey]: + checkpoint_keys = [key for key in keys if is_recurrent_checkpoint_key(key)] + if not checkpoint_keys: + return [] + with self._lock: + return [key for key in checkpoint_keys if key in self._superseded] + + def _retire_owners(self, lost: list[ObjectKey]) -> int: + """Delist the superseded checkpoints that reference ``lost`` pages.""" + if not lost: + return 0 + now = time.monotonic() + with self._lock: + generations: set[str] = set() + for key in lost: + page = self._superseded.get(key) + if page is not None: + generations.update(page.generations) + for generation in generations: + for key in self._superseded_generations.pop(generation, ()): + page = self._superseded.get(key) + if page is None: + continue + page.generations.discard(generation) + if not page.generations: + # Nothing listed needs it any more. + page.stale_at = min(page.stale_at, now) + self._continued.discard(generation) + self._stats["retired_checkpoints"] += len(generations) + retire = self._retire + if generations: + retire(sorted(generations)) + return len(generations) # ----- L2 residency --------------------------------------------------- @@ -213,17 +465,50 @@ def record_l2_absent(self, adapter_id: int, keys: list[ObjectKey]) -> None: def is_l2_resident(self, key: ObjectKey) -> bool: with self._lock: - return any(key in resident for resident in self._resident.values()) + return self._is_resident_locked(key) + + def _is_resident_locked(self, key: ObjectKey) -> bool: + return any(key in resident for resident in self._resident.values()) + + def unavailable_pages(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return the checkpoint pages that no tier holds. + + Pages recorded in L2 count as held. Of the others, a page is + unavailable when L1 does not hold it and every L2 adapter confirms + that it does not either; an adapter that cannot confirm an absence + keeps the page counted as held. + + Args: + keys: Pages of one checkpoint. + + Returns: + The unavailable pages, in the order given. + """ + checkpoint_keys = [key for key in keys if is_recurrent_checkpoint_key(key)] + with self._lock: + candidates = [ + key for key in checkpoint_keys if not self._is_resident_locked(key) + ] + if not candidates: + return [] + in_l1 = set(self._l1_present(candidates)) + candidates = [key for key in candidates if key not in in_l1] + if not candidates: + return [] + return self._l2_absent(candidates) def superseded_in_adapter( self, adapter_id: int, max_bytes: int ) -> tuple[list[ObjectKey], int]: - """Return superseded pages one adapter holds, up to ``max_bytes``.""" + """Return stale superseded pages one adapter holds, up to ``max_bytes``.""" victims: list[ObjectKey] = [] total = 0 + now = time.monotonic() with self._lock: resident: dict[ObjectKey, int] = self._resident.get(adapter_id, {}) - for key in self._superseded.items(): + for key, page in self._superseded.items(): + if page.stale_at > now: + continue size = resident.get(key) if size is None: continue @@ -248,8 +533,9 @@ def needs_persist_before_evict(self, key: ObjectKey) -> bool: """Whether L1 must keep ``key`` until its L2 write completes. True for a current checkpoint page that is not in L2 yet, while its - write is still expected to finish. A superseded page, an ordinary KV - chunk or a page whose write timed out may be evicted. + write is still expected to finish. A superseded page (stale or in its + grace period), an ordinary KV chunk or a page whose write timed out + may be evicted. """ if not self._write_on_evict or not is_recurrent_checkpoint_key(key): return False @@ -257,7 +543,7 @@ def needs_persist_before_evict(self, key: ObjectKey) -> bool: with self._lock: if key in self._superseded: return False - if any(key in resident for resident in self._resident.values()): + if self._is_resident_locked(key): self._pending.pop(key, None) return False requested = self._pending.get(key) @@ -301,10 +587,14 @@ def observations(self) -> list[tuple[int | float, dict[str, object]]]: return [(value, {"stat": name}) for name, value in status.items()] def report_status(self) -> dict: + now = time.monotonic() with self._lock: return { "write_on_evict": self._write_on_evict, "superseded_pages_tracked": len(self._superseded), + "superseded_pages_in_grace": sum( + 1 for page in self._superseded.values() if page.stale_at > now + ), "l2_resident_pages_tracked": sum( len(resident) for resident in self._resident.values() ), diff --git a/lmcache/v1/distributed/config.py b/lmcache/v1/distributed/config.py index add647874cc..2077b1c8732 100644 --- a/lmcache/v1/distributed/config.py +++ b/lmcache/v1/distributed/config.py @@ -332,13 +332,22 @@ class StorageManagerConfig: """ Total deadline for capacity-only L1 store admission retries. """ checkpoint_shutdown_flush_seconds: float = 30.0 - """ With write-on-evict checkpoint storage, time budget for writing - current checkpoint pages still in L1 to L2 during a clean shutdown. """ + """ Time budget of a clean shutdown for checkpoint work: waiting for + checkpoint stores still in flight (at most a third of it), then, with + write-on-evict checkpoint storage, writing checkpoint pages still only + in L1 to L2, current pages first. 0 skips both. """ checkpoint_write_timeout_seconds: float = 120.0 """ With write-on-evict checkpoint storage, how long L1 keeps a page whose L2 write has not completed before evicting it without a copy. """ + checkpoint_supersede_grace_seconds: float = 300.0 + """ How long a superseded checkpoint that a later prompt continued from + (a branch point, such as a turn a side request or sub-agent extended) + keeps its eviction order before its pages are dropped first. Superseded + pages are never written to L2 while serving either way. 0 drops them + first at once. """ + def __post_init__(self) -> None: if self.store_admission_timeout_seconds < 0: raise ValueError("store_admission_timeout_seconds must be >= 0") @@ -346,6 +355,8 @@ def __post_init__(self) -> None: raise ValueError("checkpoint_shutdown_flush_seconds must be >= 0") if self.checkpoint_write_timeout_seconds <= 0: raise ValueError("checkpoint_write_timeout_seconds must be > 0") + if self.checkpoint_supersede_grace_seconds < 0: + raise ValueError("checkpoint_supersede_grace_seconds must be >= 0") normalize_storage_manager_config(self) validate_storage_manager_config(self) @@ -601,9 +612,21 @@ def add_storage_manager_args( "--checkpoint-shutdown-flush-seconds", type=float, default=30.0, - help="With --l2-store-policy checkpoint_on_evict, time budget for " - "writing current recurrent checkpoints still in L1 to L2 on a clean " - "shutdown. 0 skips the flush.", + help="Time budget of a clean shutdown for recurrent checkpoints: " + "waiting for checkpoint stores still in flight (at most a third of " + "it), then, with --l2-store-policy checkpoint_on_evict, writing " + "checkpoint pages still only in L1 to L2, current ones first. " + "0 skips both.", + ) + parser.add_argument( + "--checkpoint-supersede-grace-seconds", + type=float, + default=300.0, + help="How long a superseded recurrent checkpoint that a later prompt " + "continued from (a branch point, such as a turn that a side request " + "or sub-agent extended) keeps its eviction order before its pages are " + "dropped first. Superseded pages are never written to L2 while " + "serving. 0 drops them first at once.", ) parser.add_argument( "--checkpoint-write-timeout-seconds", @@ -763,6 +786,9 @@ def parse_args_to_config( checkpoint_write_timeout_seconds=getattr( args, "checkpoint_write_timeout_seconds", 120.0 ), + checkpoint_supersede_grace_seconds=getattr( + args, "checkpoint_supersede_grace_seconds", 300.0 + ), ) return config diff --git a/lmcache/v1/distributed/l1_manager.py b/lmcache/v1/distributed/l1_manager.py index 0f4848235fd..8a147d36daa 100644 --- a/lmcache/v1/distributed/l1_manager.py +++ b/lmcache/v1/distributed/l1_manager.py @@ -1032,6 +1032,24 @@ def num_objects(self) -> int: """Return the number of objects currently tracked in L1.""" return len(self._objects) + def get_present_keys(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return the given keys that L1 holds, in the order given. + + An object counts while it is being written, since its writer may + still commit it. Like :meth:`is_key_evictable`, this does not acquire + the global L1Manager lock, so a long list never stalls other L1 + operations; each key is a point-in-time check. Nothing is read or + touched: eviction order, listeners and events are unchanged. + + Args: + keys: Keys to look up. + + Returns: + The keys with an L1 object. + """ + objects = self._objects + return [key for key in keys if key in objects] + def is_key_evictable(self, key: ObjectKey) -> bool: """Check if a key is eligible for eviction (not locked). diff --git a/lmcache/v1/distributed/l2_adapters/base.py b/lmcache/v1/distributed/l2_adapters/base.py index 1e118821611..6b55ff37904 100644 --- a/lmcache/v1/distributed/l2_adapters/base.py +++ b/lmcache/v1/distributed/l2_adapters/base.py @@ -668,6 +668,24 @@ def get_existing_key_sizes(self) -> Mapping[ObjectKey, int]: """ return _EMPTY_KEY_SIZES + def absent_keys(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return the given keys this adapter can confirm it does not hold. + + Answered synchronously from cheap local state (in-process metadata, + or one ``stat`` per key for a filesystem), so it must see objects + that other processes sharing the backend stored. Callers use it to + tell a lost object from one this process never saw. The default + confirms nothing, which keeps callers from treating an unknown key + as lost. + + Args: + keys: Keys to check. + + Returns: + The keys that are certainly absent, in the order given. + """ + return [] + def _initialize_usage(self, key_sizes: Mapping[ObjectKey, int]) -> None: """Seed byte accounting before the adapter accepts operations. diff --git a/lmcache/v1/distributed/l2_adapters/fault_inject_l2_adapter.py b/lmcache/v1/distributed/l2_adapters/fault_inject_l2_adapter.py index 9add5ba95f7..46ea8bdb811 100644 --- a/lmcache/v1/distributed/l2_adapters/fault_inject_l2_adapter.py +++ b/lmcache/v1/distributed/l2_adapters/fault_inject_l2_adapter.py @@ -402,6 +402,10 @@ def get_existing_key_sizes(self) -> Mapping[ObjectKey, int]: """Forward the persistent inventory owned by the inner adapter.""" return self._inner.get_existing_key_sizes() + def absent_keys(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Forward the absence check to the inner adapter, which holds the keys.""" + return self._inner.absent_keys(keys) + @property def supports_global_eviction(self) -> bool: """Whether the inner adapter supports aggregate usage-based eviction.""" diff --git a/lmcache/v1/distributed/l2_adapters/fs_l2_adapter.py b/lmcache/v1/distributed/l2_adapters/fs_l2_adapter.py index 4d98535e30b..8bd04a91b75 100644 --- a/lmcache/v1/distributed/l2_adapters/fs_l2_adapter.py +++ b/lmcache/v1/distributed/l2_adapters/fs_l2_adapter.py @@ -721,6 +721,23 @@ def _key_candidate_paths(self, key: ObjectKey) -> tuple[Path, ...]: self._require_representable_path(canonical) return tuple(candidates) + def absent_keys(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return keys with no object file under any accepted path. + + One ``stat`` per candidate path, so objects that other processes or + nodes sharing the directory published count as present. Keys whose + paths this filesystem cannot represent are not reported. + """ + absent = [] + for key in keys: + try: + paths = self._key_candidate_paths(key) + except (OSError, ValueError): + continue + if not any(os.path.exists(path) for path in paths): + absent.append(key) + return absent + async def _existing_key_path(self, key: ObjectKey) -> Optional[Path]: """Resolve the first existing canonical or compatible legacy path.""" for path in self._key_candidate_paths(key): diff --git a/lmcache/v1/distributed/l2_adapters/fs_native_l2_adapter.py b/lmcache/v1/distributed/l2_adapters/fs_native_l2_adapter.py index a712689a274..494fbe109d9 100644 --- a/lmcache/v1/distributed/l2_adapters/fs_native_l2_adapter.py +++ b/lmcache/v1/distributed/l2_adapters/fs_native_l2_adapter.py @@ -37,6 +37,8 @@ _BOUNDED_PATH_VERSION, _bounded_relative_path_to_object_key, _filename_to_object_key, + _object_key_to_filename, + _object_key_to_relative_path, ) logger = init_logger(__name__) @@ -290,6 +292,27 @@ def help(cls) -> str: ) +def _absent_native_fs_keys(base_path: str, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return keys with neither a canonical nor a legacy object file. + + One ``stat`` per path; a file another process published counts. Keys + whose paths cannot be derived are not reported. + """ + base = Path(base_path) + absent = [] + for key in keys: + try: + paths = ( + base / _object_key_to_relative_path(key), + base / _object_key_to_filename(key), + ) + except ValueError: + continue + if not any(os.path.exists(path) for path in paths): + absent.append(key) + return absent + + def _create_fs_native_l2_adapter( config: L2AdapterConfigBase, l1_memory_desc: "Optional[L1MemoryDesc]" = None, @@ -334,6 +357,7 @@ def _create_fs_native_l2_adapter( "read_ahead_size": config.read_ahead_size, }, initial_key_sizes=initial_key_sizes, + absence_probe=lambda keys: _absent_native_fs_keys(config.base_path, keys), ) except Exception: native_client.close() diff --git a/lmcache/v1/distributed/l2_adapters/mock_l2_adapter.py b/lmcache/v1/distributed/l2_adapters/mock_l2_adapter.py index 37d9475c863..fcedbfc3ddb 100644 --- a/lmcache/v1/distributed/l2_adapters/mock_l2_adapter.py +++ b/lmcache/v1/distributed/l2_adapters/mock_l2_adapter.py @@ -364,6 +364,11 @@ def delete(self, keys: list[ObjectKey]) -> None: if deleted_keys: self._notify_keys_deleted(deleted_keys, deleted_sizes) + def absent_keys(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return the keys the mock does not hold; it keeps objects in memory.""" + with self._lock: + return [key for key in keys if key not in self._memory_objects] + # ``get_usage()`` is inherited from ``L2AdapterInterface``, which derives # the report from the byte counters maintained by ``_notify_keys_*``. # ``_current_size_bytes`` above is a local within-batch accumulator diff --git a/lmcache/v1/distributed/l2_adapters/native_connector_l2_adapter.py b/lmcache/v1/distributed/l2_adapters/native_connector_l2_adapter.py index 2ed9fb7d106..4df19581eb2 100644 --- a/lmcache/v1/distributed/l2_adapters/native_connector_l2_adapter.py +++ b/lmcache/v1/distributed/l2_adapters/native_connector_l2_adapter.py @@ -24,7 +24,7 @@ from collections import defaultdict from dataclasses import dataclass from types import MappingProxyType -from typing import Any, Mapping +from typing import Any, Callable, Mapping import select import threading @@ -98,6 +98,11 @@ class _PendingStore: buffer_owners: list[MemoryObj] +def _confirm_no_absence(keys: list[ObjectKey]) -> list[ObjectKey]: + """Default absence probe: a backend without one cannot confirm a miss.""" + return [] + + class NativeConnectorL2Adapter(L2AdapterInterface): """ Wraps a pybind-wrapped C++ IStorageConnector to @@ -127,14 +132,20 @@ def __init__( type_name: str = "", extra_status: dict[str, Any] | None = None, initial_key_sizes: Mapping[ObjectKey, int] | None = None, + absence_probe: Callable[[list[ObjectKey]], list[ObjectKey]] = ( + _confirm_no_absence + ), ) -> None: """Create an adapter around a native storage connector. ``initial_key_sizes`` is a startup-only inventory supplied by a persistent backend. It seeds both capacity accounting and delete-time size tracking before the demultiplexer can process any completion. + ``absence_probe`` answers :meth:`absent_keys` synchronously (a + filesystem backend stats each object); the default confirms nothing. """ super().__init__(max_capacity_bytes=int(max_capacity_gb * (1024**3))) + self._absence_probe = absence_probe existing_key_sizes = dict(initial_key_sizes or {}) self._initialize_usage(existing_key_sizes) self._client = native_client @@ -206,6 +217,13 @@ def get_existing_key_sizes(self) -> Mapping[ObjectKey, int]: with self._lock: return MappingProxyType(dict(self._key_sizes)) + def absent_keys(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return the keys the backend's absence probe confirms are missing. + + Backends without a probe (the default) confirm nothing. + """ + return self._absence_probe(keys) + def has_inflight_store_for_keys(self, keys: list[ObjectKey]) -> bool: """Return whether a native store may still read any requested buffer. diff --git a/lmcache/v1/distributed/l2_adapters/serde_wrapper.py b/lmcache/v1/distributed/l2_adapters/serde_wrapper.py index a890127c18f..baff3622d32 100644 --- a/lmcache/v1/distributed/l2_adapters/serde_wrapper.py +++ b/lmcache/v1/distributed/l2_adapters/serde_wrapper.py @@ -316,6 +316,10 @@ def get_existing_key_sizes(self) -> Mapping[ObjectKey, int]: """Forward the persistent inventory owned by the inner adapter.""" return self._inner.get_existing_key_sizes() + def absent_keys(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Forward the absence check to the inner adapter, which holds the keys.""" + return self._inner.absent_keys(keys) + def delete(self, keys: list[ObjectKey]) -> None: self._inner.delete(keys) diff --git a/lmcache/v1/distributed/storage_controllers/eviction_controller.py b/lmcache/v1/distributed/storage_controllers/eviction_controller.py index 61befdda3c3..80ef8d66c9e 100644 --- a/lmcache/v1/distributed/storage_controllers/eviction_controller.py +++ b/lmcache/v1/distributed/storage_controllers/eviction_controller.py @@ -212,22 +212,38 @@ def _request_persist(self, to_persist: list[ObjectKey]) -> None: self._retention.request_persist(list(dict.fromkeys(to_persist))) def _drop_superseded(self) -> int: - """Evict superseded checkpoint pages from L1 before any LRU victim.""" + """Evict stale superseded checkpoint pages from L1 before LRU victims. + + Superseded checkpoints that lose their last page copy this way are + retired from the directory before the pages are deleted. + """ retention = self._retention if retention is None: return 0 - victims = [ - key - for key in retention.superseded_keys() - if self._l1_manager.is_key_evictable(key) - ][:_SUPERSEDED_DROP_BATCH] + victims: list[ObjectKey] = [] + for key in retention.take_droppable_superseded(): + if self._l1_manager.is_key_evictable(key): + victims.append(key) + if len(victims) >= _SUPERSEDED_DROP_BATCH: + break if not victims: return 0 + retention.retire_before_l1_delete(victims) result = self._l1_manager.delete(victims) dropped = sum(1 for error in result.values() if error == L1Error.SUCCESS) retention.record_l1_superseded_drops(dropped) return dropped + def _discard(self, keys: list[ObjectKey]) -> None: + """Delete ``keys`` from L1 without an L2 write. + + Superseded checkpoints whose last page copy this deletes are retired + from the directory first. + """ + if self._retention is not None: + self._retention.retire_before_l1_delete(keys) + self._l1_manager.delete(keys) + def request_immediate_eviction(self) -> None: """Wake the eviction loop for a capacity-blocked store.""" self._immediate_request.set() @@ -488,13 +504,13 @@ def execute_eviction_action(self, action: EvictionAction): else: logger.error("L2 eviction destination requires writeback") logger.error("Treating it as DISCARD.") - self._l1_manager.delete(action.keys) + self._discard(action.keys) elif action.destination == EvictionDestination.DISCARD: - self._l1_manager.delete(action.keys) + self._discard(action.keys) else: logger.error("Unsupported eviction destination: %s", action.destination) logger.error("Treating it as DISCARD.") - self._l1_manager.delete(action.keys) + self._discard(action.keys) def emergency_evict_bytes( self, @@ -895,7 +911,7 @@ def set_checkpoint_retention(self, retention: "CheckpointRetention") -> None: self._retention = retention def _evict_superseded(self, state: L2AdapterEvictionState) -> bool: - """Delete superseded checkpoint pages first; return True if any went.""" + """Delete stale superseded checkpoint pages first; True if any went.""" retention = self._retention if retention is None: return False @@ -907,7 +923,7 @@ def _evict_superseded(self, state: L2AdapterEvictionState) -> bool: if not victims: return False self._execute_eviction_action( - state.adapter, + state, EvictionAction(keys=victims, destination=EvictionDestination.DISCARD), ) retention.record_l2_superseded_evictions(len(victims), size) @@ -1035,7 +1051,7 @@ def _check_and_evict_global(self, state: L2AdapterEvictionState): return actions = state.eviction_policy.get_eviction_actions(eviction_ratio) for action in actions: - self._execute_eviction_action(state.adapter, action) + self._execute_eviction_action(state, action) def _check_and_evict_by_cache_salt(self, state: L2AdapterEvictionState): """Per-``cache_salt`` eviction driven by :class:`QuotaManager`. @@ -1092,19 +1108,22 @@ def _check_and_evict_by_cache_salt(self, state: L2AdapterEvictionState): for destination, keys in pending.items(): self._execute_eviction_action( - state.adapter, + state, EvictionAction(keys=keys, destination=destination), ) def _execute_eviction_action( - self, adapter: L2AdapterInterface, action: EvictionAction + self, state: L2AdapterEvictionState, action: EvictionAction ): - if action.destination == EvictionDestination.DISCARD: - adapter.delete(action.keys) - else: + adapter = state.adapter + if action.destination != EvictionDestination.DISCARD: logger.error("Unsupported eviction destination: %s", action.destination) logger.error("Treating it as DISCARD.") - adapter.delete(action.keys) + if self._retention is not None: + # Superseded checkpoints losing their last page copy are delisted + # before the pages go. + self._retention.retire_before_l2_delete(state.adapter_id, action.keys) + adapter.delete(action.keys) if action.keys: get_event_bus().publish( diff --git a/lmcache/v1/distributed/storage_manager.py b/lmcache/v1/distributed/storage_manager.py index c6b5e845933..e8dba946e1f 100644 --- a/lmcache/v1/distributed/storage_manager.py +++ b/lmcache/v1/distributed/storage_manager.py @@ -77,6 +77,9 @@ logger = init_logger(__name__) +# Seconds between progress lines of the shutdown checkpoint flush. +_FLUSH_PROGRESS_SECONDS = 2.0 + class StorageManager: def __init__(self, config: StorageManagerConfig): @@ -118,10 +121,15 @@ def __init__(self, config: StorageManagerConfig): self._checkpoint_shutdown_flush_seconds = ( config.checkpoint_shutdown_flush_seconds ) + # Monotonic deadline of the shutdown drain and flush; 0 until started. + self._shutdown_deadline = 0.0 self._checkpoint_retention = CheckpointRetention( write_on_evict=store_policy.writes_checkpoints_on_evict(), persist_timeout=config.checkpoint_write_timeout_seconds, + supersede_grace_seconds=config.checkpoint_supersede_grace_seconds, ) + self._checkpoint_retention.set_l1_lookup(self._l1_manager.get_present_keys) + self._checkpoint_retention.set_l2_absence(self._absent_from_l2) # L1 eviction controller self._eviction_controller = L1EvictionController( @@ -612,19 +620,58 @@ def _track_checkpoint_residency( adapter_id, list(inventory), list(inventory.values()) ) - def _flush_current_checkpoints(self) -> None: + def _absent_from_l2(self, keys: list[ObjectKey]) -> list[ObjectKey]: + """Return the keys every L2 adapter confirms it does not hold.""" + absent = keys + for _adapter_id, _descriptor, adapter in self._snapshot_adapters(): + if not absent: + break + try: + confirmed = set(adapter.absent_keys(absent)) + except Exception: + logger.exception("L2 absence check failed; assuming keys present") + return [] + absent = [key for key in absent if key in confirmed] + return absent + + def start_shutdown(self) -> float: + """Start the clean-shutdown budget for checkpoint work. + + The budget (``checkpoint_shutdown_flush_seconds``) covers waiting for + checkpoint stores still in flight and then writing checkpoint pages + that are only in L1 to L2. Idempotent. + + Returns: + The monotonic deadline of the budget. + """ + if not self._shutdown_deadline: + self._shutdown_deadline = ( + time.monotonic() + self._checkpoint_shutdown_flush_seconds + ) + return self._shutdown_deadline + + def _flush_current_checkpoints(self, deadline: float) -> None: """Write checkpoint pages still only in L1 to L2 on shutdown. - Current pages go first. Superseded pages follow within the same - budget: while serving they leave L1 without a write to spare flash, - but a shutdown writes each at most once, and their manifests stay - listed, so a request that continues an older line after the restart - (a branch, a resumed turn, a side request that extended the - conversation) restores instead of finding its pages gone. + Current pages go first, then superseded pages still in their grace + period, then stale superseded pages, all within the same deadline: + while serving, superseded pages leave L1 without a write to spare + flash, but a shutdown writes each at most once, and their manifests + stay listed, so a request that continues an older line after the + restart (a branch, a resumed turn, a side request that extended the + conversation) restores instead of finding its pages gone. Progress is + logged every few seconds, so a flush cut short by SIGKILL still says + how far it got. + + Args: + deadline: Monotonic time at which the flush gives up. """ retention = self._checkpoint_retention - budget = self._checkpoint_shutdown_flush_seconds - if not retention.write_on_evict or budget <= 0 or not self._has_l2_adapters(): + if ( + not retention.write_on_evict + or not self._checkpoint_shutdown_flush_seconds + or not self._has_l2_adapters() + ): return count = self._l1_manager.num_objects() keys, _ = self._l1_manager.get_evictable_keys(limit=count, scan_limit=count) @@ -633,31 +680,50 @@ def _flush_current_checkpoints(self) -> None: for key in keys if is_recurrent_checkpoint_key(key) and not retention.is_l2_resident(key) ] - current = [key for key in unwritten if not retention.is_superseded(key)] - superseded = [key for key in unwritten if retention.is_superseded(key)] - pending = current + superseded + superseded = set(retention.superseded_keys()) + current = [key for key in unwritten if key not in superseded] + stale = set(retention.take_droppable_superseded()) + in_grace = [key for key in unwritten if key in superseded and key not in stale] + old = [key for key in unwritten if key in superseded and key in stale] + pending = current + in_grace + old if not pending: return + start = time.monotonic() logger.info( "Writing %d current and %d superseded checkpoint pages to L2 before " - "shutdown (budget %.0f s)", + "shutdown (%.1f s left)", len(current), - len(superseded), - budget, + len(in_grace) + len(old), + max(0.0, deadline - start), ) - start = time.monotonic() - retention.request_persist(pending) - while time.monotonic() - start < budget: + if retention.request_persist(pending) < len(pending): + # Writes requested while serving may have failed; the store + # controller skips the ones still in flight. + self._store_controller.submit_reused_keys(pending) + last_report = start + while True: pending = [key for key in pending if not retention.is_l2_resident(key)] - if not pending: + now = time.monotonic() + if not pending or now >= deadline: break - time.sleep(0.2) + if now - last_report >= _FLUSH_PROGRESS_SECONDS: + last_report = now + logger.info( + "Shutdown checkpoint flush: %d pages left after %.1f s", + len(pending), + now - start, + ) + time.sleep(0.05) if pending: + left = set(pending) logger.warning( - "Shutdown checkpoint flush left %d pages without an L2 copy " - "after %.0f s", + "Shutdown checkpoint flush left %d pages without an L2 copy after " + "%.1f s (%d current, %d superseded); after the restart, lookups " + "retire the checkpoints that need them and fall back to shorter ones", len(pending), time.monotonic() - start, + sum(1 for key in current if key in left), + sum(1 for key in in_grace + old if key in left), ) else: logger.info( @@ -1035,6 +1101,7 @@ def delete_l1_keys( and the number refused because they were locked (non-force only). Missing keys are a no-op, so the operation is idempotent. """ + self._checkpoint_retention.retire_before_l1_delete(keys) results = self._l1_manager.delete(keys, force=force) deleted = sum(1 for err in results.values() if err == L1Error.SUCCESS) skipped = sum(1 for err in results.values() if err == L1Error.KEY_IS_LOCKED) @@ -1390,14 +1457,20 @@ def clear(self, force: bool = False): If False (default), only clear unlocked objects, keeping write-locked and read-locked objects intact. """ + retention = self._checkpoint_retention + retention.retire_before_l1_delete(retention.superseded_keys()) self._l1_manager.clear(force=force) def close(self): """ Close the storage manager and release all resources. + + With write-on-evict checkpoint storage, checkpoint pages still only + in L1 are first written to L2 within what is left of the budget that + :meth:`start_shutdown` started (it starts here if nothing did). """ try: - self._flush_current_checkpoints() + self._flush_current_checkpoints(self.start_shutdown()) except Exception: logger.exception("Shutdown checkpoint flush failed") self._l1_manager.begin_shutdown() diff --git a/lmcache/v1/multiprocess/checkpoint_index.py b/lmcache/v1/multiprocess/checkpoint_index.py index 0f3d09d54dc..c5434143060 100644 --- a/lmcache/v1/multiprocess/checkpoint_index.py +++ b/lmcache/v1/multiprocess/checkpoint_index.py @@ -511,11 +511,26 @@ def invalidate(self, generation: str) -> None: generation: Failed generation, not merely its shared token prefix. Invalidating a replaced generation cannot remove its replacement. """ + self.retire([generation]) + + def retire(self, generations: list[str]) -> None: + """Remove generations whose payload pages are lost, in one transaction. + + Args: + generations: Published generations to delist; unknown ones are + ignored. A replacement published for the same prefix is kept. + """ + if not generations: + return + rows = [(generation,) for generation in generations] with self._lock, self._db: for table in ("checkpoint_tails", "checkpoints"): - self._db.execute( - f"DELETE FROM {table} WHERE generation=?", (generation,) - ) + self._db.executemany(f"DELETE FROM {table} WHERE generation=?", rows) + + def pending_count(self) -> int: + """Return the number of staged generations not yet published.""" + with self._lock: + return len(self._pending) def report_status(self) -> dict[str, int]: """Return directory counts without asserting that payload pages are resident. diff --git a/lmcache/v1/multiprocess/http_server.py b/lmcache/v1/multiprocess/http_server.py index ed6f9b84e91..1c6f00cfca5 100644 --- a/lmcache/v1/multiprocess/http_server.py +++ b/lmcache/v1/multiprocess/http_server.py @@ -167,6 +167,9 @@ async def lifespan(app: FastAPI): # Shutdown logger.info("Shutting down LMCache HTTP server...") + # Engines stopping at the same time may still be storing checkpoints; + # serve them for a bounded time, first, before the message queue closes. + engine.drain_for_shutdown() coordinator_registration_task = getattr( app.state, "coordinator_registration_task", None ) diff --git a/lmcache/v1/multiprocess/modules/checkpoint.py b/lmcache/v1/multiprocess/modules/checkpoint.py index 4dab8a1644f..2aa06000ffc 100644 --- a/lmcache/v1/multiprocess/modules/checkpoint.py +++ b/lmcache/v1/multiprocess/modules/checkpoint.py @@ -6,6 +6,7 @@ from pathlib import Path import json import threading +import time # First Party from lmcache.logging import init_logger @@ -43,6 +44,15 @@ # L1, which holds far fewer checkpoints. _PERSISTENT_INDEX_MAX_ENTRIES = 65536 _RAM_INDEX_MAX_ENTRIES = 8192 +# Candidates a lookup may retire because no tier holds their pages any more +# before it gives up; each is one directory query and one residency check. +_MAX_LOOKUP_RETIREMENTS = 256 +# While the server stops: with no store in flight, how long to wait after the +# last store RPC for a straggler (a request finishing as the engine stops +# begins its last store at once), and with stores in flight, how long without +# any store RPC before their producer is taken as gone. +_DRAIN_SETTLE_SECONDS = 0.5 +_DRAIN_IDLE_SECONDS = 3.0 def _kind(manifest: CheckpointManifest) -> str | None: @@ -90,8 +100,10 @@ class CheckpointModule: None for 65536 with ``index_path`` and 8192 for a RAM-only directory. - Shutdown requires workers to finish or abort all submitted copy leases. - A timeout must not recycle SHM while a worker can still access its bytes. + At shutdown, :meth:`drain_stores` gives workers a bounded time to finish + or abort their copy leases while the message queue still serves them. + A lease left after that makes :meth:`close` raise; its buffers are never + recycled while a worker can still access their bytes. """ def __init__( @@ -129,6 +141,11 @@ def __init__( self._lineage_lock = threading.Lock() self._request_generations: OrderedDict[str, list[str]] = OrderedDict() self._generation_request: dict[str, str] = {} + # Monotonic time of the last store RPC, for the shutdown drain. + self._last_store_activity = 0.0 + ctx.storage_manager.checkpoint_retention.set_retirement( + self._retire_generations + ) @property def context(self) -> MPCacheServerContext: @@ -172,27 +189,68 @@ def begin(self, manifest: CheckpointManifest) -> bool: changed or already published generation raises ValueError. """ checkpoint_page_groups(manifest) + self._last_store_activity = time.monotonic() return self._index.begin(manifest) def find(self, prefixes: tuple[CheckpointPrefix, ...]) -> CheckpointManifest | None: """Return the longest complete candidate for authenticated prefix roots. A candidate is not a cache hit until all payload ranks restore it. - Every rank's payload pages are refreshed in L1 and L2 eviction order: - a resumed conversation is often served from the engine's own cache, - so its pages are neither read nor rewritten, and its oldest pages, - which every later checkpoint of the conversation needs, would - otherwise keep the recency of their first write. + A listed candidate whose pages no tier holds any more (dropped, + evicted or deleted since it was published, or not written before a + restart) is retired from the directory, and the next shorter one is + tried at once. This needs L2 adapters whose inventory is complete; + otherwise the restore validates the pages. + + Every rank's payload pages of the returned candidate are refreshed in + L1 and L2 eviction order: a resumed conversation is often served from + the engine's own cache, so its pages are neither read nor rewritten, + and its oldest pages, which every later checkpoint of the + conversation needs, would otherwise keep the recency of their first + write. A superseded candidate becomes current again: a request is + continuing from it. """ - manifest = self._index.find(prefixes) - if manifest is not None: + storage = self._ctx.storage_manager + retention = storage.checkpoint_retention + retired: list[tuple[int, int, int]] = [] + found: CheckpointManifest | None = None + for _ in range(_MAX_LOOKUP_RETIREMENTS + 1): + manifest = self._index.find(prefixes) + if manifest is None: + break try: keys = _payload_keys(manifest) except ValueError: # Retrieval rejects the same manifest and invalidates it. - return manifest - self._ctx.storage_manager.touch_keys(keys) - return manifest + found = manifest + break + lost = retention.unavailable_pages(keys) + if not lost: + storage.touch_keys(keys) + if retention.mark_generation_current(manifest.generation, keys): + logger.debug( + "Superseded checkpoint of %d tokens was found again; " + "it is current", + manifest.prefix.num_tokens, + ) + found = manifest + break + self._index.invalidate(manifest.generation) + retired.append((manifest.prefix.num_tokens, len(lost), len(keys))) + if retired: + tokens, lost_pages, pages = retired[0] + logger.info( + "Checkpoint lookup retired %d listed checkpoints whose pages no " + "tier holds (longest %d tokens, %d of %d pages lost); %s", + len(retired), + tokens, + lost_pages, + pages, + f"found one of {found.prefix.num_tokens} tokens" + if found is not None + else "found none", + ) + return found def supersede( self, prefixes: tuple[CheckpointPrefix, ...], generation: str, request: str @@ -201,10 +259,14 @@ def supersede( Called after ``generation`` is published by ``request``, with the roots of the token sequence that produced it. Only a ``prompt`` - checkpoint supersedes: the complete prompt of a new request replaces - every published shorter checkpoint of its sequence that an earlier - request produced, together with the other checkpoints of those - requests that it has moved past (a response endpoint that the chat + checkpoint supersedes. Its longest published ancestor from an earlier + request is the checkpoint the new prompt continues from; it stays + current, because the new request may be a side request (a title or + summary request, a sub-agent fork) while the conversation goes on + from the same point later. The other published ancestors of the + sequence from earlier requests are superseded: the conversation has + moved past them twice. So are the other checkpoints of their requests + that the new prompt has moved past (a response endpoint that the chat template rewrote, a prefill tail). A checkpoint of such a request that is at least as long as the new prompt is kept: the new request branched off before it (a retry, or an aborted turn resumed with @@ -213,9 +275,10 @@ def supersede( which other conversations share, are kept. Pages the new checkpoint references are current again. - Superseded pages are evicted first and are never written to L2 on - eviction. Their manifests stay listed, so a request that branches - from an older turn can still restore one while its pages last. + Superseded pages are never written to L2 while serving. They are + evicted first, at once, or after the grace period when a later prompt + had continued from their checkpoint. When a superseded page's last + copy is deleted, its checkpoints are retired from the directory. Returns: Number of pages newly marked superseded. @@ -231,20 +294,29 @@ def supersede( retention.mark_current(current_keys) if _kind(current) != "prompt": return 0 - victims: dict[str, CheckpointManifest] = {} - for ancestor in self._index.ancestors(prefixes, current.prefix.num_tokens): - owner = self._request_of(ancestor.generation) - if owner == request: - continue - victims[ancestor.generation] = ancestor - for sibling in self._generations_of(owner): - if sibling not in victims and sibling != generation: - manifest = self._index.get(sibling) - if ( - manifest is not None - and manifest.prefix.num_tokens < current.prefix.num_tokens - ): - victims[sibling] = manifest + ancestors = [ + ancestor + for ancestor in self._index.ancestors(prefixes, current.prefix.num_tokens) + if self._request_of(ancestor.generation) != request + ] + if not ancestors: + return 0 + parent = ancestors[-1] + retention.mark_continued(parent.generation) + retention.mark_generation_current(parent.generation, _payload_keys(parent)) + victims: dict[str, CheckpointManifest] = { + ancestor.generation: ancestor for ancestor in ancestors[:-1] + } + for ancestor in ancestors: + for sibling in self._generations_of(self._request_of(ancestor.generation)): + if sibling in victims or sibling in (generation, parent.generation): + continue + manifest = self._index.get(sibling) + if ( + manifest is not None + and manifest.prefix.num_tokens < current.prefix.num_tokens + ): + victims[sibling] = manifest selected = sorted( ( victim @@ -262,10 +334,11 @@ def supersede( if selected: logger.debug( "Prompt checkpoint of %d tokens superseded %d older checkpoints " - "(%d pages)", + "(%d pages); it continues from one of %d tokens", current.prefix.num_tokens, len(selected), marked, + parent.prefix.num_tokens, ) return marked @@ -293,6 +366,7 @@ def _generations_of(self, request: str | None) -> list[str]: def abort(self, generation: str) -> bool: """Prevent publication; rank copy leases still require explicit finish.""" + self._last_store_activity = time.monotonic() self._index.abort(generation) return True @@ -303,6 +377,7 @@ def prepare_store( Invalid layouts, ranks or duplicate active rank stores raise ValueError. """ + self._last_store_activity = time.monotonic() admission = self._payloads.prepare_store(manifest, rank) if admission is AdmissionFailure.BUSY: return CheckpointLeaseResponse("busy") @@ -314,7 +389,11 @@ def prepare_store( def finish_store(self, lease_id: str, success: bool) -> bool: """Finish drained D2H work; True means all ranks published the manifest.""" - return self._payloads.finish_store(lease_id, success) + self._last_store_activity = time.monotonic() + try: + return self._payloads.finish_store(lease_id, success) + finally: + self._last_store_activity = time.monotonic() def begin_retrieve( self, manifest: CheckpointManifest, rank: int @@ -364,13 +443,91 @@ def report_status(self) -> dict[str, dict[str, int]]: } } + def drain_stores(self, deadline: float) -> None: + """Keep serving checkpoint stores until they settle or ``deadline``. + + Call it when the server starts to stop, while its message queue still + serves requests. An engine stopping at the same time may still be + copying its last checkpoints; each rank that finishes publishes them, + so the shutdown flush can write them to L2. + + Returns once no store lease or unpublished generation remains and no + store RPC arrived for a short settle time, once no store RPC arrived + for a few seconds while stores are still in flight (their producer is + gone), or at ``deadline``. New stores are admitted meanwhile. + + Args: + deadline: Monotonic time at which to stop waiting. + """ + start = time.monotonic() + announced = False + while True: + now = time.monotonic() + in_flight = ( + self._payloads.report_status()["store_leases"] + + self._index.pending_count() + ) + quiet = now - self._last_store_activity + if ( + (not in_flight and quiet >= _DRAIN_SETTLE_SECONDS) + or quiet >= _DRAIN_IDLE_SECONDS + or now >= deadline + ): + break + if not announced: + announced = True + logger.info( + "Serving checkpoint stores for up to %.1f s before shutdown " + "(%d in flight)", + max(0.0, deadline - now), + in_flight, + ) + time.sleep(0.02) + if in_flight: + logger.warning( + "%d checkpoint stores were still in flight after %.1f s of " + "shutdown; those checkpoints stay unpublished", + in_flight, + time.monotonic() - start, + ) + elif announced: + logger.info( + "Checkpoint stores settled after %.1f s of shutdown", + time.monotonic() - start, + ) + + def _retire_generations(self, generations: list[str]) -> None: + """Delist superseded checkpoints whose last page copy is being deleted.""" + try: + self._index.retire(generations) + except Exception: + # The directory may already be closed during shutdown; a listed + # checkpoint without pages is retired by the next lookup instead. + logger.debug("Could not retire %d checkpoints", len(generations)) + return + logger.debug( + "Retired %d superseded checkpoints whose last pages left the cache", + len(generations), + ) + def close(self) -> None: """Close the directory after copy leases drain, otherwise raise RuntimeError. This module never frees buffers solely because a copy took too long. - The owning process shutdown must coordinate GPU worker termination. + The owning process shutdown must coordinate GPU worker termination: + :meth:`drain_stores` gives workers a bounded time first, and + ``MPCacheServer.close`` logs this error and still closes the storage + manager, so its shutdown flush runs and its shared memory is released + without reusing a lease's buffers. """ status = self._payloads.report_status() if status["store_leases"] or status["retrieve_leases"]: - raise RuntimeError("Checkpoint worker copy leases must drain before close") + raise RuntimeError( + f"Checkpoint worker copy leases must drain before close " + f"({status['store_leases']} store, " + f"{status['retrieve_leases']} retrieve)" + ) + self._ctx.storage_manager.checkpoint_retention.clear_retirement( + self._retire_generations + ) self._index.close() diff --git a/lmcache/v1/multiprocess/server.py b/lmcache/v1/multiprocess/server.py index 15e7f661df8..6a478a26a0d 100644 --- a/lmcache/v1/multiprocess/server.py +++ b/lmcache/v1/multiprocess/server.py @@ -79,6 +79,12 @@ logger = init_logger(__name__) +# Upper bound of the shutdown drain of checkpoint stores in flight. It takes +# at most a third of the checkpoint shutdown budget, so the L2 flush after it +# keeps most of the budget, and a stop with Docker's default 10 s timeout +# still leaves the flush time. +_MAX_SHUTDOWN_DRAIN_SECONDS = 5.0 + def _unlink_configured_l1_shm(shm_name: str) -> None: """Remove the exact named L1 pool before measuring available capacity. @@ -164,10 +170,37 @@ def report_status(self) -> dict: status.update(module.report_status()) return status + def drain_for_shutdown(self) -> None: + """Let checkpoint stores still in flight finish before shutdown. + + Call it once shutdown starts and before the message queue server + closes, so an engine stopping at the same time can still publish the + checkpoints it is copying. It starts the storage manager's shutdown + budget and uses at most a third of it (and at most a few seconds); + :meth:`close` then writes checkpoint pages to L2 within the rest. + """ + deadline = self._context.storage_manager.start_shutdown() + now = time.monotonic() + drain_deadline = now + min( + _MAX_SHUTDOWN_DRAIN_SECONDS, max(0.0, deadline - now) / 3 + ) + for module in self._modules: + if isinstance(module, CheckpointModule): + module.drain_stores(drain_deadline) + def close(self) -> None: - """Close all modules and release shared resources.""" + """Close all modules and release shared resources. + + A module that fails to close is logged and skipped, so the storage + manager still flushes checkpoints and releases the shared memory. + """ for module in self._modules: - module.close() + try: + module.close() + except Exception: + logger.exception( + "Closing %s failed; continuing shutdown", type(module).__name__ + ) self._context.close() logger.info("MPCacheServer closed") @@ -541,6 +574,7 @@ def run_cache_server( time.sleep(1) except KeyboardInterrupt: logger.info("Shutting down server...") + engine.drain_for_shutdown() event_bus.stop() server.close() engine.close() diff --git a/tests/v1/distributed/test_checkpoint_retention.py b/tests/v1/distributed/test_checkpoint_retention.py index 8459e902b77..8d10a6660e2 100644 --- a/tests/v1/distributed/test_checkpoint_retention.py +++ b/tests/v1/distributed/test_checkpoint_retention.py @@ -152,3 +152,100 @@ def test_observations_report_every_stat_and_l2_checkpoint_bytes(): assert stats["l2_checkpoint_bytes"] == 150 assert stats["l2_resident_pages_tracked"] == 2 assert "write_on_evict" not in stats + + +def test_a_continued_checkpoint_keeps_its_order_for_the_grace_period(): + """Pages of a checkpoint that a later prompt continued from (a possible + branch point) are dropped first only after the grace period; pages of any + other superseded checkpoint at once. Neither is ever held for a write.""" + retention = CheckpointRetention(write_on_evict=True, supersede_grace_seconds=60.0) + fork_point, passed = make_key(1), make_key(2) + retention.mark_continued("fork-point") + with patch("time.monotonic", return_value=100.0): + retention.mark_superseded("fork-point", [fork_point]) + retention.mark_superseded("passed", [passed]) + assert retention.take_droppable_superseded() == [passed] + status = retention.report_status() + assert status["superseded_pages_in_grace"] == 1 + assert not retention.needs_persist_before_evict(fork_point) + with patch("time.monotonic", return_value=161.0): + assert retention.take_droppable_superseded() == [fork_point, passed] + + +def test_found_checkpoint_is_current_again_and_can_be_superseded_anew(): + retention = CheckpointRetention() + keys = [make_key(1), make_key(2)] + retention.mark_superseded("gen", keys) + assert retention.mark_generation_current("gen", keys) + assert not retention.is_superseded_generation("gen") + assert retention.superseded_keys() == [] + assert not retention.mark_generation_current("gen", keys) + assert retention.mark_superseded("gen", keys) == 2 + assert retention.report_status()["restored_checkpoints"] == 1 + + +def test_superseded_pages_no_tier_holds_are_forgotten(): + present: list[ObjectKey] = [] + retention = CheckpointRetention() + retention.set_l1_lookup(lambda keys: [key for key in keys if key in present]) + in_l1, in_l2, gone = make_key(1), make_key(2), make_key(3) + present.append(in_l1) + retention.record_l2_present(0, [in_l2], [10]) + retention.mark_superseded("gen", [in_l1, in_l2, gone]) + assert retention.take_droppable_superseded() == [in_l1] + assert retention.superseded_keys() == [in_l1, in_l2] + + +def test_last_copy_deletions_retire_their_superseded_checkpoints_first(): + """L1 loses the last copy of a page without an L2 copy; an adapter loses + it when neither L1 nor another adapter holds it. Only then are the + superseded checkpoints that reference it retired, once.""" + retired: list[list[str]] = [] + in_l1: list[ObjectKey] = [] + retention = CheckpointRetention() + retention.set_l1_lookup(lambda keys: [key for key in keys if key in in_l1]) + retention.set_retirement(retired.append) + written, unwritten, shared = make_key(1), make_key(2), make_key(3) + retention.record_l2_present(0, [written], [10]) + retention.mark_superseded("a", [written, shared]) + retention.mark_superseded("b", [unwritten, shared]) + + assert retention.retire_before_l1_delete([written]) == 0 + assert retention.retire_before_l1_delete([unwritten, make_key(9, False)]) == 1 + assert retired == [["b"]] + assert not retention.is_superseded_generation("b") + # "a" still needs its written page; L1 holds a copy of the shared one. + in_l1.append(shared) + retention.record_l2_present(1, [written], [10]) + assert retention.retire_before_l2_delete(0, [written, shared]) == 0 + retention.record_l2_absent(1, [written]) + assert retention.retire_before_l2_delete(0, [written]) == 1 + assert retired == [["b"], ["a"]] + assert retention.report_status()["retired_checkpoints"] == 2 + + retention.clear_retirement(retired.append) + retention.mark_superseded("c", [make_key(4)]) + assert retention.retire_before_l1_delete([make_key(4)]) == 1 + assert retired == [["b"], ["a"]] + + +def test_unavailable_pages_need_every_tier_to_confirm_the_absence(): + """A page is lost only when L1 lacks it, no adapter records it, and every + adapter confirms it does not hold it; without adapters L1 decides.""" + retention = CheckpointRetention() + retention.set_l1_lookup(lambda keys: [k for k in keys if k == make_key(1)]) + keys = [make_key(1), make_key(2), make_key(3), make_key(4, False)] + assert retention.unavailable_pages(keys) == [make_key(2), make_key(3)] + asked: list[list[ObjectKey]] = [] + + def adapters_confirm(candidates: list[ObjectKey]) -> list[ObjectKey]: + asked.append(candidates) + return [key for key in candidates if key == make_key(3)] + + retention.set_l2_absence(adapters_confirm) + retention.record_l2_present(0, [make_key(2)], [10]) + assert retention.unavailable_pages(keys) == [make_key(3)] + # Only pages neither L1 nor the L2 records hold are checked. + assert asked == [[make_key(3)]] + assert retention.unavailable_pages([make_key(1), make_key(2)]) == [] + assert asked == [[make_key(3)]] diff --git a/tests/v1/distributed/test_fs_l2_adapter_persistence.py b/tests/v1/distributed/test_fs_l2_adapter_persistence.py index 9d2710ef35b..b20dfe1cbc4 100644 --- a/tests/v1/distributed/test_fs_l2_adapter_persistence.py +++ b/tests/v1/distributed/test_fs_l2_adapter_persistence.py @@ -392,3 +392,20 @@ def test_rejects_filesystem_below_protocol_limit(tmp_path: Path) -> None: pytest.raises(ValueError, match=r"PC_NAME_MAX >= 255, got 254"), ): FSL2Adapter(FSL2AdapterConfig(base_path=str(tmp_path))) + + +def test_absent_keys_sees_objects_other_adapters_stored(tmp_path: Path) -> None: + """Absence is checked on disk: an object another adapter instance stored + in the shared directory is present, a deleted one is absent.""" + stored, missing = _long_key(salt_suffix="a"), _long_key(salt_suffix="b") + reader = FSL2Adapter(FSL2AdapterConfig(base_path=str(tmp_path))) + writer = FSL2Adapter(FSL2AdapterConfig(base_path=str(tmp_path))) + try: + assert reader.absent_keys([stored, missing]) == [stored, missing] + _wait_for_store(writer, stored, b"payload") + assert reader.absent_keys([stored, missing]) == [missing] + writer.delete([stored]) + assert reader.absent_keys([stored, missing]) == [stored, missing] + finally: + writer.close() + reader.close() diff --git a/tests/v1/distributed/test_fs_native_startup_scan.py b/tests/v1/distributed/test_fs_native_startup_scan.py index 4637bc530d4..71792c5c804 100644 --- a/tests/v1/distributed/test_fs_native_startup_scan.py +++ b/tests/v1/distributed/test_fs_native_startup_scan.py @@ -15,6 +15,7 @@ _object_key_to_relative_path, ) from lmcache.v1.distributed.l2_adapters.fs_native_l2_adapter import ( + _absent_native_fs_keys, _scan_existing_key_sizes, ) @@ -193,3 +194,22 @@ def test_scan_rejects_duplicate_decoded_keys(tmp_path) -> None: with pytest.raises(RuntimeError, match="multiple filenames"): _scan_existing_key_sizes(str(tmp_path)) + + +def test_absence_probe_sees_objects_any_process_published(tmp_path) -> None: + """An object file under its canonical or legacy name counts as present, + whoever wrote it; only keys with neither file are confirmed absent.""" + canonical, legacy, missing = _key(1), _key(2), _key(3) + oversized = ObjectKey( + chunk_hash=ObjectKey.IntHash2Bytes(4), + model_name="m" * 240, + kv_rank=0, + object_group_id=1, + ) + for key in (canonical, oversized): + path = tmp_path / _object_key_to_relative_path(key) + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(b"object") + (tmp_path / _object_key_to_filename(legacy)).write_bytes(b"object") + keys = [canonical, legacy, missing, oversized] + assert _absent_native_fs_keys(str(tmp_path), keys) == [missing] diff --git a/tests/v1/distributed/test_mock_l2_adapter.py b/tests/v1/distributed/test_mock_l2_adapter.py index fa9ee71090a..39cbed5af43 100644 --- a/tests/v1/distributed/test_mock_l2_adapter.py +++ b/tests/v1/distributed/test_mock_l2_adapter.py @@ -886,3 +886,13 @@ def test_multiple_listeners_all_notified(self, adapter): assert len(l1.stored) == 1 assert len(l2.stored) == 1 + + +def test_absent_keys_confirms_only_objects_the_mock_does_not_hold(adapter): + """The mock holds objects in memory, so it can confirm an absence.""" + stored, missing = create_object_key(1), create_object_key(2) + assert adapter.absent_keys([stored, missing]) == [stored, missing] + _store_and_wait(adapter, stored, create_memory_obj()) + assert adapter.absent_keys([stored, missing]) == [missing] + adapter.delete([stored]) + assert adapter.absent_keys([stored, missing]) == [stored, missing] diff --git a/tests/v1/multiprocess/test_checkpoint_shutdown_drain.py b/tests/v1/multiprocess/test_checkpoint_shutdown_drain.py new file mode 100644 index 00000000000..4e4d387bf41 --- /dev/null +++ b/tests/v1/multiprocess/test_checkpoint_shutdown_drain.py @@ -0,0 +1,243 @@ +# SPDX-License-Identifier: Apache-2.0 +"""A stopping server finishes checkpoint stores in flight, then writes them. + +When a container stops, the engine and the cache server get SIGTERM at the +same time. The engine may still be copying its last checkpoints; the server +keeps serving those stores for a bounded time before its message queue +closes, and its shutdown flush then writes them to L2. +""" + +# Standard +from collections.abc import Iterator +from contextlib import contextmanager +from mmap import mmap +from pathlib import Path +from types import SimpleNamespace +from typing import cast +from unittest.mock import MagicMock, call, patch +import asyncio +import threading +import time +import uuid + +# First Party +from lmcache.v1.distributed.config import ( + EvictionConfig, + L1ManagerConfig, + L1MemoryManagerConfig, + StorageManagerConfig, +) +from lmcache.v1.distributed.l2_adapters.config import L2AdaptersConfig +from lmcache.v1.distributed.l2_adapters.fs_native_l2_adapter import ( + FSNativeL2AdapterConfig, +) +from lmcache.v1.distributed.storage_manager import StorageManager +from lmcache.v1.multiprocess import http_server +from lmcache.v1.multiprocess.checkpoint_index import CheckpointManifest +from lmcache.v1.multiprocess.engine_context import MPCacheServerContext +from lmcache.v1.multiprocess.modules.checkpoint import CheckpointModule +from lmcache.v1.multiprocess.posix_shm import shm_open_pool_as_mmap +from lmcache.v1.multiprocess.server import MPCacheServer +from tests.v1.multiprocess.test_checkpoint_side_requests import ( + POOL_BYTES, + checkpoint, + content, + page_keys, + prefix, + publish, + restores, +) + + +@contextmanager +def open_server( + path: Path, *, budget_seconds: float = 30.0 +) -> Iterator[tuple[MPCacheServer, CheckpointModule, mmap]]: + """A cache server with a checkpoint module over RAM and native-FS L2. + + The caller shuts it down with ``drain_for_shutdown`` and ``close``. + """ + name = f"lmcache_l1_pool_checkpoint_drain_{uuid.uuid4().hex}" + storage = StorageManager( + StorageManagerConfig( + L1ManagerConfig(L1MemoryManagerConfig(POOL_BYTES, False, shm_name=name)), + EvictionConfig(eviction_policy="LRU"), + l2_adapter_config=L2AdaptersConfig( + [FSNativeL2AdapterConfig(str(path / "payloads"))] + ), + store_policy="checkpoint_on_evict", + checkpoint_shutdown_flush_seconds=budget_seconds, + ) + ) + context = SimpleNamespace( + storage_manager=storage, + shm_pool_info={"shm_name": name, "pool_size": POOL_BYTES}, + close=storage.close, + ) + module = CheckpointModule( + cast(MPCacheServerContext, context), path / "directory.sqlite3" + ) + server = MPCacheServer(cast(MPCacheServerContext, context), [module]) + with shm_open_pool_as_mmap(name, POOL_BYTES) as mapping: + yield server, module, mapping + + +def copy_pages( + module: CheckpointModule, mapping: mmap, entry: CheckpointManifest +) -> str: + """Admit a store and copy its pages, as a worker does before finishing.""" + assert module.begin(entry) + lease = module.prepare_store(entry, 0) + assert lease.status == "ready" + slots = [slot for group in lease.slots for slot in group] + for key, (offset, size) in zip(page_keys(entry), slots, strict=True): + if offset >= 0: + mapping[offset : offset + size] = content(key) + return lease.lease_id + + +def restores_after_restart(path: Path, entry: CheckpointManifest) -> bool: + with open_server(path) as (server, module, mapping): + try: + return module.find((entry.prefix,)) == entry and restores( + module, mapping, entry + ) + finally: + server.close() + + +TOKENS = tuple(range(200, 213)) + + +def test_shutdown_serves_a_store_in_flight_and_writes_it(tmp_path: Path) -> None: + """The engine finishes its copy 0.4 s into the shutdown; the drain keeps + the message queue's handlers serving until then, the checkpoint is + published, and the flush writes it for the next start. Before, the + server closed at once, the lease never finished, the checkpoint module + refused to close and the storage manager never flushed.""" + entry = checkpoint(TOKENS, "response") + with open_server(tmp_path) as (server, module, mapping): + lease_id = copy_pages(module, mapping, entry) + finisher = threading.Timer(0.4, module.finish_store, (lease_id, True)) + finisher.start() + started = time.monotonic() + server.drain_for_shutdown() + elapsed = time.monotonic() - started + finisher.join() + assert 0.4 <= elapsed < 3.0 + assert module.find((entry.prefix,)) == entry + server.close() + assert restores_after_restart(tmp_path, entry) + + +def test_shutdown_drain_is_bounded_when_a_worker_never_finishes( + tmp_path: Path, +) -> None: + """A worker killed mid-copy never finishes its lease. The drain gives up + within a third of the budget, closing still flushes what was published, + and the unfinished checkpoint is never listed.""" + published = checkpoint(TOKENS, "response") + stuck = checkpoint(TOKENS + (1, 2, 3), "prompt") + with open_server(tmp_path, budget_seconds=1.5) as (server, module, mapping): + publish(module, mapping, published, "a") + copy_pages(module, mapping, stuck) + started = time.monotonic() + server.drain_for_shutdown() + assert time.monotonic() - started < 1.0 + server.close() + assert restores_after_restart(tmp_path, published) + with open_server(tmp_path) as (server, module, _mapping): + try: + assert module.find((prefix(stuck.prefix.tail_tokens),)) == published + finally: + server.close() + + +def test_shutdown_drain_returns_at_once_when_idle(tmp_path: Path) -> None: + with open_server(tmp_path) as (server, _module, _mapping): + started = time.monotonic() + server.drain_for_shutdown() + assert time.monotonic() - started < 0.3 + server.close() + + +def test_shutdown_drain_is_bounded_while_stores_keep_arriving( + tmp_path: Path, +) -> None: + """An engine that keeps storing cannot hold the shutdown: the drain ends + after a third of a 3 s budget and the flush keeps the rest.""" + stop = threading.Event() + with open_server(tmp_path, budget_seconds=3.0) as (server, module, mapping): + + def keep_storing() -> None: + # The same prefix again and again: every store is admitted at + # once, since its pages are already in RAM. + turn = 0 + while not stop.is_set(): + turn += 1 + publish(module, mapping, checkpoint(TOKENS[:4], "prompt"), f"r{turn}") + time.sleep(0.05) + + producer = threading.Thread(target=keep_storing) + producer.start() + try: + time.sleep(0.2) + started = time.monotonic() + server.drain_for_shutdown() + elapsed = time.monotonic() - started + finally: + stop.set() + producer.join() + assert 0.8 <= elapsed < 1.5 + server.close() + + +def test_server_close_still_closes_storage_after_a_module_fails() -> None: + """A module that refuses to close must not skip the storage manager's + close, which runs the checkpoint flush and releases shared memory.""" + failing, other, context = MagicMock(), MagicMock(), MagicMock() + failing.close.side_effect = RuntimeError("leases must drain") + MPCacheServer(context, [failing, other]).close() + other.close.assert_called_once_with() + context.close.assert_called_once_with() + + +def test_http_lifespan_drains_before_the_message_queue_closes() -> None: + """The HTTP server's shutdown serves stores in flight before it closes + the message queue, then closes the engine.""" + order = MagicMock() + zmq_server, engine = order.zmq_server, order.engine + configs = { + "mp": SimpleNamespace(runtime_plugin_config=SimpleNamespace(locations=[])), + "storage_manager": MagicMock(), + "observability": MagicMock(), + } + + async def run() -> None: + async with http_server.lifespan(MagicMock()): + pass + + with ( + patch.dict(http_server._configs, configs, clear=True), + patch.object( + http_server, "run_cache_server", return_value=(zmq_server, engine) + ), + patch.object(http_server, "build_context", return_value=MagicMock()), + patch.object(http_server, "get_event_bus", return_value=MagicMock()), + ): + asyncio.run(run()) + shutdown = [ + entry + for entry in order.mock_calls + if entry + in ( + call.engine.drain_for_shutdown(), + call.zmq_server.close(), + call.engine.close(), + ) + ] + assert shutdown == [ + call.engine.drain_for_shutdown(), + call.zmq_server.close(), + call.engine.close(), + ] diff --git a/tests/v1/multiprocess/test_checkpoint_side_requests.py b/tests/v1/multiprocess/test_checkpoint_side_requests.py new file mode 100644 index 00000000000..e247ab1f329 --- /dev/null +++ b/tests/v1/multiprocess/test_checkpoint_side_requests.py @@ -0,0 +1,415 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Side requests, branch points and lost pages of recurrent checkpoints. + +A conversation's latest checkpoint can be extended by a request that is not +the conversation's next turn: a title or summary request, a sub-agent fork. +Supersession must keep what the conversation still continues from, and a +checkpoint whose pages are gone must stop being offered instead of failing a +restore. +""" + +# Standard +from collections.abc import Callable, Iterator +from contextlib import ExitStack, contextmanager +from mmap import mmap +from pathlib import Path +from types import SimpleNamespace +from typing import cast +import hashlib +import json +import math +import time +import uuid + +# Third Party +import pytest + +# First Party +from lmcache.v1.distributed.api import ObjectKey +from lmcache.v1.distributed.config import ( + EvictionConfig, + L1ManagerConfig, + L1MemoryManagerConfig, + StorageManagerConfig, +) +from lmcache.v1.distributed.l2_adapters.config import L2AdaptersConfig +from lmcache.v1.distributed.l2_adapters.fs_native_l2_adapter import ( + FSNativeL2AdapterConfig, +) +from lmcache.v1.distributed.storage_manager import StorageManager +from lmcache.v1.multiprocess.checkpoint_index import ( + CheckpointManifest, + CheckpointPrefix, +) +from lmcache.v1.multiprocess.checkpoint_storage import checkpoint_object_keys +from lmcache.v1.multiprocess.engine_context import MPCacheServerContext +from lmcache.v1.multiprocess.modules.checkpoint import CheckpointModule +from lmcache.v1.multiprocess.posix_shm import shm_open_pool_as_mmap + +POOL_BYTES = 4 * 1024 * 1024 +PAGE_BYTES = 64 * 1024 +# Tokens per attention page. +BLOCK = 4 +NAMESPACE = "weights-and-layout-and-salt" + + +def digest(text: str) -> str: + return hashlib.sha256(text.encode()).hexdigest() + + +def prefix(tokens: tuple[int, ...]) -> CheckpointPrefix: + return CheckpointPrefix(NAMESPACE, 4096, b"a" * 32, tokens) + + +def checkpoint(tokens: tuple[int, ...], kind: str) -> CheckpointManifest: + """A one-rank bundle shaped like a request-boundary checkpoint. + + Attention pages are keyed by the tokens they cover, so a conversation's + checkpoints share their full pages; the partial last page, the recurrent + state and the auxiliary page belong to this endpoint alone. + """ + positions = list(range(math.ceil(len(tokens) / BLOCK))) + attention = [ + digest(f"attention:{tokens[: min((position + 1) * BLOCK, len(tokens))]}") + for position in positions + ] + return CheckpointManifest( + uuid.uuid4().hex, + prefix(tokens), + 1, + json.dumps( + { + "schema_version": 2, + "kind": kind, + "page_groups": [ + { + "name": "target.attention.0", + "page_bytes": PAGE_BYTES, + "positions": positions, + "content_keys": attention, + }, + { + "name": "target.recurrent.0", + "page_bytes": PAGE_BYTES, + "positions": [0], + "content_keys": [digest(f"recurrent:{tokens}")], + }, + { + "name": "target-draft-auxiliary", + "page_bytes": PAGE_BYTES, + "positions": [0], + "content_keys": [digest(f"auxiliary:{kind}:{tokens}")], + }, + ], + } + ).encode(), + ) + + +def page_keys(entry: CheckpointManifest) -> list[ObjectKey]: + return [key for group in checkpoint_object_keys(entry, 0) for key in group] + + +def content(key: ObjectKey) -> bytes: + """The bytes a page holds, derived from its content key.""" + return bytes([key.chunk_hash[0]]) * PAGE_BYTES + + +@contextmanager +def open_module( + path: Path, + *, + grace_seconds: float = 300.0, + flush_seconds: float = 0.0, +) -> Iterator[tuple[CheckpointModule, StorageManager, mmap]]: + """A checkpoint module over a 4 MiB L1 and a native-filesystem L2 whose + inventory is complete, with checkpoint pages written on L1 eviction.""" + name = f"lmcache_l1_pool_checkpoint_side_{uuid.uuid4().hex}" + storage = StorageManager( + StorageManagerConfig( + L1ManagerConfig(L1MemoryManagerConfig(POOL_BYTES, False, shm_name=name)), + EvictionConfig(eviction_policy="LRU"), + l2_adapter_config=L2AdaptersConfig( + [FSNativeL2AdapterConfig(str(path / "payloads"))] + ), + store_policy="checkpoint_on_evict", + checkpoint_shutdown_flush_seconds=flush_seconds, + checkpoint_supersede_grace_seconds=grace_seconds, + ) + ) + with ExitStack() as cleanup: + cleanup.callback(storage.close) + mapping = cleanup.enter_context(shm_open_pool_as_mmap(name, POOL_BYTES)) + module = CheckpointModule( + cast( + MPCacheServerContext, + SimpleNamespace( + storage_manager=storage, + shm_pool_info={"shm_name": name, "pool_size": POOL_BYTES}, + ), + ), + path / "directory.sqlite3", + ) + cleanup.callback(module.close) + yield module, storage, mapping + + +def publish( + module: CheckpointModule, mapping: mmap, entry: CheckpointManifest, request: str +) -> None: + """Store, publish and report one checkpoint as the vLLM bridge does.""" + assert module.begin(entry) + lease = module.prepare_store(entry, 0) + assert lease.status == "ready", lease.status + slots = [slot for group in lease.slots for slot in group] + for key, (offset, size) in zip(page_keys(entry), slots, strict=True): + if offset >= 0: + mapping[offset : offset + size] = content(key) + assert module.finish_store(lease.lease_id, True) + module.supersede((entry.prefix,), entry.generation, request) + + +def restores( + module: CheckpointModule, mapping: mmap, entry: CheckpointManifest +) -> bool: + """Whether every page of ``entry`` restores with the bytes it published.""" + lease = module.begin_retrieve(entry, 0) + if lease.status != "pending": + return False + deadline = time.monotonic() + 10 + while lease.status == "pending" and time.monotonic() < deadline: + lease = module.poll_retrieve(lease.lease_id) + time.sleep(0.001) + if lease.status != "ready": + return False + try: + slots = [slot for group in lease.slots for slot in group] + for key, (offset, size) in zip(page_keys(entry), slots, strict=True): + assert mapping[offset : offset + size] == content(key) + finally: + module.finish_retrieve(lease.lease_id) + return True + + +def cycle_l1(module: CheckpointModule, mapping: mmap, conversations: int) -> None: + """Publish other conversations until RAM has been refilled about twice.""" + for conversation in range(conversations): + tokens = tuple(1000 * (conversation + 2) + i for i in range(4 * BLOCK)) + publish(module, mapping, checkpoint(tokens, "prompt"), f"filler-{conversation}") + + +def wait_for( + condition: Callable[[], bool], message: str, timeout: float = 10.0 +) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if condition(): + return + time.sleep(0.02) + raise AssertionError(message) + + +# A conversation of two turns: prompt, then the response endpoint. +BASE = tuple(range(100, 100 + 2 * BLOCK)) +PROMPT = checkpoint(BASE, "prompt") +RESPONSE_TOKENS = BASE + (7, 8, 9, 10, 11) + + +def test_side_request_keeps_the_turn_it_extended_under_ram_pressure( + tmp_path: Path, +) -> None: + """A title request extends the conversation's response; the user then + continues the conversation after RAM pressure evicted everything. + + The side request continues from the response, so the response stays + current: when RAM evicts it, it is written to L2 like any current page, + and the conversation's next turn restores it. Before, the side request + superseded it, its state pages were dropped without a write, and the + listed checkpoint failed its restore ('K of M pages were readable'). + """ + response = checkpoint(RESPONSE_TOKENS, "response") + with open_module(tmp_path) as (module, storage, mapping): + publish(module, mapping, checkpoint(BASE, "prompt"), "a") + publish(module, mapping, response, "a") + title = checkpoint(RESPONSE_TOKENS + (50, 51, 52), "prompt") + publish(module, mapping, title, "title") + publish( + module, + mapping, + checkpoint(title.prefix.tail_tokens + (53,), "response"), + "title", + ) + retention = storage.checkpoint_retention + cycle_l1(module, mapping, 16) + wait_for( + lambda: retention.is_l2_resident(page_keys(response)[0]), + "RAM pressure never reached the conversation", + ) + + # The conversation continues from its response. + assert module.find((prefix(RESPONSE_TOKENS + (60, 61)),)) == response + assert restores(module, mapping, response) + # It stayed current: its pages left RAM with a write, none was dropped. + assert not any(retention.is_superseded(key) for key in page_keys(response)) + assert retention.report_status()["retired_checkpoints"] >= 1 + + +def test_a_superseded_checkpoint_is_retired_before_its_last_pages_go( + tmp_path: Path, +) -> None: + """Pages of a checkpoint the conversation moved past twice are dropped + without a write; the checkpoint is delisted before they go, so a lookup + misses it cleanly and falls back to a shorter one at once.""" + instruction = checkpoint(BASE[:BLOCK], "instruction") + response = checkpoint(RESPONSE_TOKENS, "response") + turn_2 = checkpoint(RESPONSE_TOKENS + (20, 21), "prompt") + turn_3 = checkpoint(turn_2.prefix.tail_tokens + (22, 23, 24), "prompt") + with open_module(tmp_path, grace_seconds=0.0) as (module, storage, mapping): + publish(module, mapping, instruction, "system") + publish(module, mapping, PROMPT, "a") + publish(module, mapping, response, "a") + publish(module, mapping, turn_2, "b") + publish(module, mapping, turn_3, "c") + retention = storage.checkpoint_retention + # Turn 3 moved past turn 1's prompt and response; turn 2 is where it + # continued from and stays current. + assert all(retention.is_superseded(key) for key in page_keys(PROMPT)[-2:]) + assert all(retention.is_superseded(key) for key in page_keys(response)[-2:]) + assert not any(retention.is_superseded(key) for key in page_keys(turn_2)) + + cycle_l1(module, mapping, 16) + wait_for( + lambda: retention.report_status()["retired_checkpoints"] >= 2, + "superseded checkpoints were not retired when their pages were dropped", + ) + assert not any( + retention.is_l2_resident(key) for key in page_keys(response)[-2:] + ) + # A branch from the end of turn 1 misses both turn-1 checkpoints at + # once and gets the shared instruction. + assert module.find((prefix(RESPONSE_TOKENS + (99,)),)) == instruction + assert restores(module, mapping, instruction) + # The conversation itself still continues from turn 3. + assert module.find((prefix(turn_3.prefix.tail_tokens + (30,)),)) == turn_3 + assert restores(module, mapping, turn_3) + + +@pytest.mark.parametrize("grace", [300.0, 0.0]) +def test_a_fork_point_is_not_dropped_first_during_the_grace_period( + tmp_path: Path, grace: float +) -> None: + """A sub-agent forks from the conversation's response and runs two turns, + which supersedes the response. As a branch point it is not dropped first + during the grace period; when the conversation finds it again it is + current again and restores. With no grace, RAM pressure drops it first, + without a write, and it is retired: the lookup misses it cleanly.""" + response = checkpoint(RESPONSE_TOKENS, "response") + fork_prompt = checkpoint(RESPONSE_TOKENS + (40, 41), "prompt") + fork_response = checkpoint(fork_prompt.prefix.tail_tokens + (42,), "response") + fork_turn_2 = checkpoint(fork_response.prefix.tail_tokens + (43, 44), "prompt") + continuation = prefix(RESPONSE_TOKENS + (60,)) + with open_module(tmp_path, grace_seconds=grace) as (module, storage, mapping): + retention = storage.checkpoint_retention + publish(module, mapping, PROMPT, "a") + publish(module, mapping, response, "a") + publish(module, mapping, fork_prompt, "fork") + publish(module, mapping, fork_response, "fork") + publish(module, mapping, fork_turn_2, "fork-2") + state = page_keys(response)[-3:] + assert all(retention.is_superseded(key) for key in state) + droppable = set(retention.take_droppable_superseded()) + # The prompt and the fork's prompt were passed, not continued from. + assert set(page_keys(PROMPT)[-2:]) <= droppable + assert set(page_keys(fork_prompt)[-2:]) <= droppable + if grace: + assert not droppable & set(state) + assert module.find((continuation,)) == response + assert not any(retention.is_superseded(key) for key in state) + assert restores(module, mapping, response) + else: + assert set(state) <= droppable + cycle_l1(module, mapping, 16) + wait_for( + lambda: module.find((continuation,)) is None, + "the dropped fork point stayed listed", + ) + assert not any(retention.is_l2_resident(key) for key in state) + + +def test_lookup_retires_checkpoints_whose_pages_no_tier_holds(tmp_path: Path) -> None: + """A current page lost outside supersession (an admin delete of a page + that never reached L2) makes its checkpoint unrestorable. The lookup + retires it and returns the next shorter complete checkpoint in the same + call instead of offering a restore that fails.""" + response = checkpoint(RESPONSE_TOKENS, "response") + turn_2 = checkpoint(RESPONSE_TOKENS + (20, 21), "prompt") + with open_module(tmp_path) as (module, storage, mapping): + publish(module, mapping, PROMPT, "a") + publish(module, mapping, response, "a") + publish(module, mapping, turn_2, "b") + continuation = prefix(turn_2.prefix.tail_tokens + (30,)) + assert module.find((continuation,)) == turn_2 + deleted, _ = storage.delete_l1_keys(page_keys(turn_2)[-1:]) + assert deleted == 1 + assert module.find((continuation,)) == response + assert restores(module, mapping, response) + # Retired, not merely skipped: an exact lookup no longer lists it. + assert module.find((turn_2.prefix,)) == response + + +def test_restart_retires_checkpoints_that_were_not_written(tmp_path: Path) -> None: + """After a restart the directory lists checkpoints whose pages stayed in + RAM. A lookup skips and retires them and restores the newest one that + reached L2, without a failed restore first.""" + response = checkpoint(RESPONSE_TOKENS, "response") + turn_2 = checkpoint(RESPONSE_TOKENS + (20, 21), "prompt") + with open_module(tmp_path, flush_seconds=30.0) as (module, _storage, mapping): + publish(module, mapping, PROMPT, "a") + publish(module, mapping, response, "a") + with open_module(tmp_path, flush_seconds=0.0) as (module, _storage, mapping): + publish(module, mapping, turn_2, "b") + with open_module(tmp_path) as (module, _storage, mapping): + found = module.find((prefix(turn_2.prefix.tail_tokens + (30,)),)) + assert found == response + assert restores(module, mapping, response) + assert module.find((turn_2.prefix,)) == response + + +def test_ordinary_deletions_never_reach_the_directory() -> None: + """Deleting KV chunks or current checkpoint pages retires nothing and + takes no directory lock: only superseded pages carry owners.""" + name = f"lmcache_l1_pool_checkpoint_fastpath_{uuid.uuid4().hex}" + storage = StorageManager( + StorageManagerConfig( + L1ManagerConfig(L1MemoryManagerConfig(POOL_BYTES, False, shm_name=name)), + EvictionConfig(eviction_policy="LRU"), + ) + ) + retired: list[list[str]] = [] + try: + retention = storage.checkpoint_retention + retention.set_retirement(retired.append) + ordinary = [ + ObjectKey(ObjectKey.IntHash2Bytes(i), "some-model", 0) for i in range(3) + ] + current = page_keys(PROMPT) + assert retention.retire_before_l1_delete(ordinary + current) == 0 + assert retention.retire_before_l2_delete(0, ordinary + current) == 0 + retention.mark_superseded("old", current[-1:]) + assert retention.retire_before_l1_delete(ordinary + current) == 1 + assert retired == [["old"]] + finally: + storage.close() + + +def test_lookup_trusts_pages_another_process_wrote(tmp_path: Path) -> None: + """Two servers share one filesystem tier and directory. A checkpoint the + other one wrote after this one started is not in this one's records, but + its files exist: the lookup keeps it and it restores. Only pages that no + tier holds make a lookup retire a checkpoint.""" + response = checkpoint(RESPONSE_TOKENS, "response") + with open_module(tmp_path) as (module, _storage, mapping): + with open_module(tmp_path, flush_seconds=30.0) as (writer, _, writer_map): + publish(writer, writer_map, response, "a") + assert module.find((prefix(RESPONSE_TOKENS + (60,)),)) == response + assert restores(module, mapping, response) diff --git a/tests/v1/multiprocess/test_checkpoint_storage.py b/tests/v1/multiprocess/test_checkpoint_storage.py index 5a24023213b..d65c5758211 100644 --- a/tests/v1/multiprocess/test_checkpoint_storage.py +++ b/tests/v1/multiprocess/test_checkpoint_storage.py @@ -1555,7 +1555,8 @@ def test_index_lists_ancestors_of_a_sequence_without_touching_lru() -> None: def test_supersede_marks_only_pages_older_turns_alone_reference() -> None: - """A new prompt supersedes earlier requests' checkpoints, not its own.""" + """A new prompt supersedes earlier requests' checkpoints, not its own, and + not the one it continues from until a later prompt moves past that too.""" name = f"lmcache_l1_pool_checkpoint_supersede_{uuid.uuid4().hex}" with open_store(shm_name=name) as (_, _, storage, mapping): module = CheckpointModule( @@ -1576,8 +1577,9 @@ def test_supersede_marks_only_pages_older_turns_alone_reference() -> None: tail_2 = conversation_manifest((1, 2, 3), "t2", kind="prefill_tail") prompt_2 = conversation_manifest((1, 2, 3, 4), "p2", kind="prompt") response_2 = conversation_manifest((1, 2, 3, 4, 6), "r2") + prompt_3 = conversation_manifest((1, 2, 3, 4, 7, 8), "p3", kind="prompt") turns = (instruction, prompt_1, response_1, tail_2, prompt_2, response_2) - for turn in turns: + for turn in (*turns, prompt_3): # Shared pages already resident come back without a slot. assert module._index.begin(turn) lease = module._payloads.prepare_store(turn, 0) @@ -1593,6 +1595,9 @@ def test_supersede_marks_only_pages_older_turns_alone_reference() -> None: def roots(entry: CheckpointManifest) -> tuple[CheckpointPrefix, ...]: return (entry.prefix,) + def unique(*entries: CheckpointManifest) -> set[ObjectKey]: + return {key for entry in entries for key in entry_keys(entry)} + retention = storage.checkpoint_retention requests = ("r0", "r1", "r1", "r2", "r2", "r2") # Only a prompt supersedes, and turn 1 has no earlier request. @@ -1600,14 +1605,13 @@ def roots(entry: CheckpointManifest) -> tuple[CheckpointPrefix, ...]: assert module.supersede(roots(turn), turn.generation, request) == 0 assert retention.superseded_keys() == [] + # Turn 2 continues from turn 1's prompt, which stays current; it + # has moved past the rewritten response. marked = module.supersede(roots(prompt_2), prompt_2.generation, "r2") - current = set(entry_keys(prompt_2)) - expected = (set(entry_keys(prompt_1)) | set(entry_keys(response_1))) - ( - current - ) + expected = unique(response_1) - unique(prompt_2) assert marked == len(expected) > 0 assert set(retention.superseded_keys()) == expected - for kept in (instruction, tail_2, prompt_2): + for kept in (instruction, prompt_1, tail_2, prompt_2): assert not any(retention.is_superseded(k) for k in entry_keys(kept)) # The same request's response supersedes nothing. assert module.supersede(roots(response_2), response_2.generation, "r2") == 0 @@ -1618,13 +1622,24 @@ def roots(entry: CheckpointManifest) -> tuple[CheckpointPrefix, ...]: prompt_2.generation, "r2", ) + + # Turn 3 continues from turn 2's prompt: turn 1's prompt, turn 2's + # tail and its rewritten response are passed now. + marked = module.supersede(roots(prompt_3), prompt_3.generation, "r3") + newly = unique(prompt_1, tail_2, response_2) - unique(prompt_3) + assert marked == len(newly) > 0 + assert set(retention.superseded_keys()) == expected | newly + for kept in (instruction, prompt_2, prompt_3): + assert not any(retention.is_superseded(k) for k in entry_keys(kept)) finally: module.close() def test_a_branch_keeps_the_checkpoints_it_branched_before() -> None: - """A retry or resumed turn supersedes the prompt it extends, but not the - longer response the original line continues from (lmcache-supersession-repro).""" + """A retry or resumed turn supersedes neither the prompt it continues from + nor the longer response the original line continues from + (lmcache-supersession-repro); the original line's next turn passes the + prompt, and a later one the response.""" name = f"lmcache_l1_pool_checkpoint_branch_{uuid.uuid4().hex}" with open_store(shm_name=name) as (_, _, storage, mapping): module = CheckpointModule( @@ -1642,9 +1657,12 @@ def test_a_branch_keeps_the_checkpoints_it_branched_before() -> None: # An aborted turn resumed with one other token: shorter than A's # response, so it branched before it. prompt_b = conversation_manifest((1, 2, 9), "pb", kind="prompt") - # The original line continues past A's response. + # The original line continues past A's response, then past C. prompt_c = conversation_manifest((1, 2, 5, 6, 7, 8), "pc", kind="prompt") - turns = (prompt_a, response_a, prompt_b, prompt_c) + prompt_d = conversation_manifest( + (1, 2, 5, 6, 7, 8, 10), "pd", kind="prompt" + ) + turns = (prompt_a, response_a, prompt_b, prompt_c, prompt_d) for turn in turns: assert module._index.begin(turn) lease = module._payloads.prepare_store(turn, 0) @@ -1662,16 +1680,19 @@ def test_a_branch_keeps_the_checkpoints_it_branched_before() -> None: module.supersede((response_a.prefix,), response_a.generation, "a") == 0 ) - module.supersede((prompt_b.prefix,), prompt_b.generation, "b") - unique_a = set(entry_keys(prompt_a)) - set(entry_keys(prompt_b)) + assert module.supersede((prompt_b.prefix,), prompt_b.generation, "b") == 0 + assert retention.superseded_keys() == [] + + module.supersede((prompt_c.prefix,), prompt_c.generation, "c") + unique_a = set(entry_keys(prompt_a)) - set(entry_keys(prompt_c)) assert unique_a and set(retention.superseded_keys()) == unique_a assert not any(retention.is_superseded(k) for k in entry_keys(response_a)) - module.supersede((prompt_c.prefix,), prompt_c.generation, "c") - current = set(entry_keys(prompt_c)) - assert set(retention.superseded_keys()) == ( - unique_a | (set(entry_keys(response_a)) - current) + module.supersede((prompt_d.prefix,), prompt_d.generation, "d") + assert set(retention.superseded_keys()) == unique_a | ( + set(entry_keys(response_a)) - set(entry_keys(prompt_d)) ) + assert not any(retention.is_superseded(k) for k in entry_keys(prompt_b)) finally: module.close() diff --git a/tests/v1/test_vllm_semantic_checkpoint_transfer.py b/tests/v1/test_vllm_semantic_checkpoint_transfer.py index 28e924643b3..211205ab25c 100644 --- a/tests/v1/test_vllm_semantic_checkpoint_transfer.py +++ b/tests/v1/test_vllm_semantic_checkpoint_transfer.py @@ -799,8 +799,14 @@ def drain_predecessor() -> None: worker.close() -def test_failed_restore_falls_back_to_the_longest_remaining_checkpoint() -> None: - """A restore whose pages are gone retries the directory, not the prompt.""" +@pytest.mark.parametrize("lookup_sees_loss", [True, False]) +def test_failed_restore_falls_back_to_the_longest_remaining_checkpoint( + lookup_sees_loss: bool, +) -> None: + """A checkpoint whose pages are gone falls back to the longest remaining + one, not the prompt. When the lookup can tell that no tier holds the + pages, it retires the checkpoint and answers with the shorter one at + once; otherwise the failed restore retries the directory.""" with open_checkpoint_rpc() as (client, module, mapping, _name): manager = make_manager() bridge = CheckpointSchedulerBridge( @@ -891,7 +897,16 @@ def run(task: CheckpointEngineTask) -> dict[int, bool]: consumer = make_request("consumer") attempts: list[tuple[int, bool]] = [] deadline = time.monotonic() + 10 - with patch.object(checkpoint_scheduler.logger, "info") as info: + retention = storage.checkpoint_retention + with ( + patch.object(checkpoint_scheduler.logger, "info") as info, + patch.object( + retention, + "unavailable_pages", + side_effect=None if lookup_sees_loss else lambda keys: [], + wraps=retention.unavailable_pages, + ), + ): while not bridge.poll_prefix(consumer): assert time.monotonic() < deadline for task in bridge.take_tasks(): @@ -901,12 +916,17 @@ def run(task: CheckpointEngineTask) -> dict[int, bool]: ) time.sleep(0.001) - assert attempts == [(11, False), (8, True)] - failure = next( - call.args for call in info.call_args_list if "failed" in call.args[0] - ) - assert failure[1:4] == (11, consumer.request_id, [0, 1, 2, 3]) - assert 0 <= failure[4] < 10 + if lookup_sees_loss: + assert attempts == [(8, True)] + else: + assert attempts == [(11, False), (8, True)] + failure = next( + call.args + for call in info.call_args_list + if "failed" in call.args[0] + ) + assert failure[1:4] == (11, consumer.request_id, [0, 1, 2, 3]) + assert 0 <= failure[4] < 10 assert manager.get_computed_blocks(consumer)[1] == 8 assert bridge.external_tokens(consumer) == 8 assert not bridge.has_pending