diff --git a/python/sglang/srt/mem_cache/rust_tree_core/adapter.py b/python/sglang/srt/mem_cache/rust_tree_core/adapter.py index 51bd03178..071bb9134 100644 --- a/python/sglang/srt/mem_cache/rust_tree_core/adapter.py +++ b/python/sglang/srt/mem_cache/rust_tree_core/adapter.py @@ -679,15 +679,11 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface): @property def enable_external_cache_linker(self) -> bool: - return False + return self._binding.enable_external_cache_linker() @enable_external_cache_linker.setter def enable_external_cache_linker(self, value: bool) -> None: - # TODO(Jialin): Port external cache linker support from #37091 and #37151. - if value: - raise ValueError( - "External cache linker is not supported by the Rust TreeCore" - ) + self._binding.set_enable_external_cache_linker(value) def insert_host( self, @@ -913,6 +909,27 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface): def finish_load_back(self, anchor_node_id: NodeId) -> None: self._binding.finish_load_back(anchor_node_id) + def build_external_linker_offload_transfers( + self, node_id: NodeId + ) -> Optional[list[PoolTransfer]]: + transfers = self._binding.build_external_linker_offload_transfers(node_id) + if transfers is None: + return None + return [_transfer_from_binding(transfer) for transfer in transfers] + + def mark_external_cache_stored_path( + self, from_node_id: NodeId, until_node_id: NodeId + ) -> None: + self._binding.mark_external_cache_stored_path(from_node_id, until_node_id) + + def mark_external_linker_offload_pending(self, node_id: NodeId) -> None: + self._binding.mark_external_linker_offload_pending(node_id) + + def finish_external_linker_offload( + self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool + ) -> None: + self._binding.finish_external_linker_offload(list(node_ids), ack_id, success) + @property def write_back_duplicate_reclaim_digest(self) -> int: return self._binding.write_back_duplicate_reclaim_digest() diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py b/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py index 1dfc6b999..5c163a676 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_cache_linker.py @@ -36,6 +36,7 @@ from sglang.srt.mem_cache.hicache_storage import ( ) from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.unified_cache.components import ( + ComponentType, ExternalLinkerLoadPhase, LinkerTransferPhase, TreeComponent, @@ -48,6 +49,14 @@ if TYPE_CHECKING: from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache +_EXTERNAL_LINKER_SUPPORTED_COMPONENTS = frozenset( + { + ComponentType.FULL, + ComponentType.SWA, + } +) + + class UnifiedCacheLinker(ABC): """External KV store reached directly from the device pools.""" @@ -139,6 +148,16 @@ class UnifiedCacheLinkerWrapper: cache: UnifiedRadixCache, cache_linker: UnifiedCacheLinker, ): + unsupported = set(cache.tree_components) - _EXTERNAL_LINKER_SUPPORTED_COMPONENTS + if unsupported: + names = ", ".join( + component.name for component in sorted(unsupported, key=int) + ) + raise ValueError( + "External cache linker supports only Full and SWA tree " + f"components; unsupported: {names}" + ) + self.cache = cache self.cache_linker = cache_linker # rid -> what match found, consumed by the next init_load_back. @@ -162,6 +181,7 @@ class UnifiedCacheLinkerWrapper: def match(self, key: RadixKey, req: Req, result: MatchResult) -> MatchResult: cache = self.cache + key, _ = key.maybe_to_bigram_view(cache.tree_core.is_eagle) page = cache.page_size device_hit_len = int(result.device_indices.numel()) if device_hit_len >= len(key): @@ -343,10 +363,9 @@ class UnifiedCacheLinkerWrapper: self._queue_load(req.rid, insert_result.last_device_node, load_transfers) - node = cache.resolve_node_handle(insert_result.last_device_node) - while node.id != req.last_node: - node.external_cache_stored = True - node = node.parent + cache.tree_core.mark_external_cache_stored_path( + insert_result.last_device_node, req.last_node + ) return canonical_tail, insert_result.last_device_node def _queue_load( @@ -461,20 +480,14 @@ class UnifiedCacheLinkerWrapper: def offload_nodes(self, node_ids: Sequence[NodeId]) -> None: """Persist a write-through chain, skipping nodes already in the store.""" for node_id in node_ids: - if not self.cache.resolve_node_handle(node_id).external_cache_stored: - self._offload_node(node_id) - - def _offload_node(self, node_id: NodeId) -> None: - cache = self.cache - node = cache.resolve_node_handle(node_id) - transfers = [] - for component in cache._components_tuple: - transfer = component.build_external_linker_transfer( - LinkerTransferPhase.OFFLOAD, node, None + transfers = self.cache.tree_core.build_external_linker_offload_transfers( + node_id ) - if transfer is not None: - transfers.append(transfer) + if transfers is not None: + self._offload_node(node_id, transfers) + def _offload_node(self, node_id: NodeId, transfers: list[PoolTransfer]) -> None: + cache = self.cache lock_params = cache.inc_lock_ref(node_id).to_dec_params() try: queued = self.cache_linker.offload(transfers) @@ -485,8 +498,7 @@ class UnifiedCacheLinkerWrapper: cache.dec_lock_ref(node_id, lock_params) return - cache.tree_core.mark_write_through_pending([node_id], ack_id=node_id) - node.external_cache_stored = True + cache.tree_core.mark_external_linker_offload_pending(node_id) self.pending_offloads.append(_PendingOffload(node_id, lock_params, [node_id])) def replace_pending_offload_node( @@ -528,11 +540,9 @@ class UnifiedCacheLinkerWrapper: assert len(successes) <= len(self.pending_offloads) for success in successes: pending = self.pending_offloads.pop(0) - for node_id in pending.publish_node_ids: - node = self.cache.resolve_node_handle(node_id) - if node.write_through_pending_id == pending.lock_node_id: - node.write_through_pending_id = None - node.external_cache_stored = success + self.cache.tree_core.finish_external_linker_offload( + pending.publish_node_ids, pending.lock_node_id, success + ) self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params) def start_layer_wise_loading(self) -> int: @@ -550,11 +560,9 @@ class UnifiedCacheLinkerWrapper: self.cache.dec_lock_ref(node_id, lock_params) self.pending_loads.clear() for pending in self.pending_offloads: - for node_id in pending.publish_node_ids: - node = self.cache.resolve_node_handle(node_id) - if node.write_through_pending_id == pending.lock_node_id: - node.write_through_pending_id = None - node.external_cache_stored = False + self.cache.tree_core.finish_external_linker_offload( + pending.publish_node_ids, pending.lock_node_id, False + ) self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params) self.pending_offloads.clear() diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py index 49cf6c87e..83206f4b2 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py @@ -60,6 +60,7 @@ from sglang.srt.mem_cache.unified_cache.components import ( ComponentData, ComponentType, EvictLayer, + LinkerTransferPhase, LRURefreshPhase, TreeComponent, get_and_increase_time_counter, @@ -954,7 +955,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): if self.enable_external_cache_linker: return ( - not node.external_cache_stored + self._needs_external_linker_offload(node) and node.hit_count >= self.write_through_threshold ) @@ -964,6 +965,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): and node.hit_count >= self.write_through_threshold ) + @staticmethod + def _needs_external_linker_offload(node: UnifiedTreeNode) -> bool: + """Whether neither a confirmed nor an in-flight external copy exists.""" + return not node.external_cache_stored and node.write_through_pending_id is None + def begin_insert(self, params: InsertParams) -> InsertStepResult: """Start the insert, running to its first barrier or completion.""" # Insert walks are single-flight; a live walk means re-entrancy. @@ -2140,7 +2146,12 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): while ( ancestor is not None and ancestor is not self.root_node - and not (ancestor.backuped or ancestor.external_cache_stored) + and not ancestor.backuped + and not ancestor.external_cache_stored + and ( + not self.enable_external_cache_linker + or ancestor.write_through_pending_id is None + ) ): chain.append(ancestor) ancestor = ancestor.parent @@ -2279,6 +2290,69 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface): node = node.parent return depth + def build_external_linker_offload_transfers( + self, node_id: NodeId + ) -> Optional[list[PoolTransfer]]: + """Build transfers for a node with no stored or pending external copy.""" + node = self.node_by_id(node_id) + if not self._needs_external_linker_offload(node): + return None + + transfers = [] + for component in self.components: + transfer = component.build_external_linker_transfer( + LinkerTransferPhase.OFFLOAD, node, None + ) + if transfer is not None: + transfers.append(transfer) + return transfers + + def mark_external_cache_stored_path( + self, from_node_id: NodeId, until_node_id: NodeId + ) -> None: + """Mark an externally restored path, excluding its existing anchor.""" + until_node = self.node_by_id(until_node_id) + node = self.node_by_id(from_node_id) + path = [] + while node is not until_node: + if node.parent is None: + raise RuntimeError( + f"node {until_node_id} is not an ancestor of node {from_node_id}" + ) + path.append(node) + node = node.parent + + for node in path: + node.external_cache_stored = True + + def mark_external_linker_offload_pending(self, node_id: NodeId) -> None: + """Publish an accepted external offload as pending.""" + node = self.node_by_id(node_id) + if not self._needs_external_linker_offload(node): + raise AssertionError( + f"invalid external offload state for node {node_id}: " + f"stored={node.external_cache_stored}, " + f"pending={node.write_through_pending_id}" + ) + node.write_through_pending_id = node_id + + def finish_external_linker_offload( + self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool + ) -> None: + """Finalize external-store state for an offload and its split fragments.""" + nodes = [self.node_by_id(node_id) for node_id in node_ids] + for node_id, node in zip(node_ids, nodes): + if node.write_through_pending_id != ack_id: + raise AssertionError( + f"invalid external offload state for node {node_id}: " + f"expected pending={ack_id}; got " + f"stored={node.external_cache_stored}, " + f"pending={node.write_through_pending_id}" + ) + for node in nodes: + node.write_through_pending_id = None + node.external_cache_stored |= success + def finish_write_through(self, node_ids: list[NodeId], ack_id: int) -> None: """Clear the write-through-pending mark (when it matches ack_id) and record the host store event for each acked node.""" diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py index 4c07baa43..94fc9e5ec 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py @@ -148,6 +148,7 @@ class UnifiedTreeCoreInterface(ABC): device: torch.device enable_hicache: bool enable_storage: bool + enable_external_cache_linker: bool write_through_threshold: int is_write_back: bool has_swa_host_pool: bool @@ -533,6 +534,41 @@ class UnifiedTreeCoreInterface(ABC): """Clear the in-flight H->D marks on the anchor's root path at ack time.""" ... + # ==== External Cache Linker ==== + + @abstractmethod + def build_external_linker_offload_transfers( + self, node_id: NodeId + ) -> Optional[list[PoolTransfer]]: + """Build direct device-to-external-store transfers for an eligible node. + + Return None when the node is stored externally or has an offload pending. + """ + ... + + @abstractmethod + def mark_external_cache_stored_path( + self, from_node_id: NodeId, until_node_id: NodeId + ) -> None: + """Mark the path from ``from_node_id`` to, but excluding, ``until_node_id``.""" + ... + + @abstractmethod + def mark_external_linker_offload_pending(self, node_id: NodeId) -> None: + """Publish an accepted external offload as pending.""" + ... + + @abstractmethod + def finish_external_linker_offload( + self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool + ) -> None: + """Finish one external offload for every current fragment of its node. + + A successful write confirms external storage. A failed redundant write + preserves storage already confirmed independently by a concurrent load. + """ + ... + # Order-sensitive digest of write_back duplicate-reclaim victim ids, # cross-checked across TP ranks; cores that never reclaim keep 0. write_back_duplicate_reclaim_digest: int = 0 diff --git a/rust/sglang-radix-tree/src/components/full.rs b/rust/sglang-radix-tree/src/components/full.rs index 9a4ba06d8..cd2ecfc41 100644 --- a/rust/sglang-radix-tree/src/components/full.rs +++ b/rust/sglang-radix-tree/src/components/full.rs @@ -427,6 +427,25 @@ impl TreeComponent for FullComponent { }) } + fn build_external_linker_offload_transfer( + &self, + tree_core: &UnifiedTreeCore, + node_id: NodeIdx_, + ) -> Option { + let node = tree_core.arena.node(node_id); + let keys = node + .hash_value + .as_ref() + .filter(|hashes| !hashes.is_empty())?; + let device_indices = node.try_device_value(FULL)?; + Some(PoolTransfer { + name: PoolName::Kv, + device_indices: Some(device_indices.to_kind(Kind::Int64)), + keys: Some(keys.clone()), + ..Default::default() + }) + } + fn commit_hicache_transfer( &self, tree_core: &mut UnifiedTreeCore, diff --git a/rust/sglang-radix-tree/src/components/mod.rs b/rust/sglang-radix-tree/src/components/mod.rs index da8ab956a..f098aeacf 100644 --- a/rust/sglang-radix-tree/src/components/mod.rs +++ b/rust/sglang-radix-tree/src/components/mod.rs @@ -392,6 +392,15 @@ pub trait TreeComponent { unimplemented!("TreeComponent.build_hicache_transfers") } + /// Build this component's direct device-to-external-store transfer for a node. + fn build_external_linker_offload_transfer( + &self, + _tree_core: &UnifiedTreeCore, + _node_id: NodeIdx_, + ) -> Option { + None + } + /// Post-transfer bookkeeping: store host indices, update LRU, etc. fn commit_hicache_transfer( &self, diff --git a/rust/sglang-radix-tree/src/components/swa.rs b/rust/sglang-radix-tree/src/components/swa.rs index deac00f16..262853daf 100644 --- a/rust/sglang-radix-tree/src/components/swa.rs +++ b/rust/sglang-radix-tree/src/components/swa.rs @@ -951,6 +951,35 @@ impl TreeComponent for SwaComponent { }) } + fn build_external_linker_offload_transfer( + &self, + tree_core: &UnifiedTreeCore, + node_id: NodeIdx_, + ) -> Option { + let node = tree_core.arena.node(node_id); + let hashes = node + .hash_value + .as_ref() + .filter(|hashes| !hashes.is_empty())?; + let value = node.try_device_value(SWA)?; + let num_pages = value.size()[0] as usize / tree_core.page_size; + if num_pages == 0 { + return None; + } + let num_tokens = num_pages * tree_core.page_size; + Some(PoolTransfer { + name: PoolName::Swa, + device_indices: Some( + value + .narrow(0, value.size()[0] - num_tokens as i64, num_tokens as i64) + .to_kind(Kind::Int64), + ), + keys: Some(hashes[hashes.len().saturating_sub(num_pages)..].to_vec()), + hit_policy: PoolHitPolicy::TrailingPages, + ..Default::default() + }) + } + fn commit_hicache_transfer( &self, tree_core: &mut UnifiedTreeCore, diff --git a/rust/sglang-radix-tree/src/node.rs b/rust/sglang-radix-tree/src/node.rs index f94a25298..db3cde23d 100644 --- a/rust/sglang-radix-tree/src/node.rs +++ b/rust/sglang-radix-tree/src/node.rs @@ -175,6 +175,8 @@ pub struct Node { /// Per-page hash chain; None when the node was never hashed. /// TODO: Store raw digests and hex-encode only at the Python or storage boundary. pub hash_value: Option>, + /// Whether this node is available through the direct external-cache linker. + pub external_cache_stored: bool, /// The in-flight write-through backup's ack id. pub write_through_pending_id: Option, /// Load-back anchor currently reading this node's host slots. @@ -396,6 +398,7 @@ impl Node { swa_uuid: None, swa_host_uuid: None, hash_value: Some(Vec::new()), + external_cache_stored: false, write_through_pending_id: None, load_back_pending_id: None, last_access_counter: 0, @@ -418,6 +421,7 @@ impl Node { swa_uuid: None, swa_host_uuid: None, hash_value: None, + external_cache_stored: false, write_through_pending_id: None, load_back_pending_id: None, last_access_counter: 0, @@ -753,6 +757,24 @@ pub enum TreeCoreRuntimeError { #[cfg(any(test, feature = "inspection"))] #[error("{0}")] InspectionAssertion(String), + /// Direct external-cache linking does not support this tree component. + #[error("external cache linker does not support component {component_type:?}")] + ExternalCacheLinkerUnsupportedComponent { component_type: ComponentType }, + /// The existing device anchor must be on the restored endpoint's root path. + #[error("node {until_node_id} is not an ancestor of node {from_node_id}")] + ExternalCachePathNotAncestor { + from_node_id: NodeId, + until_node_id: NodeId, + }, + /// External offload lifecycle calls must observe valid state transitions. + #[error( + "invalid external offload state for node {node_id}: stored={stored}, pending={pending_id:?}" + )] + InvalidExternalCacheOffloadState { + node_id: NodeId, + stored: bool, + pending_id: Option, + }, } // Unigram and bigram child keys. diff --git a/rust/sglang-radix-tree/src/python_bindings.rs b/rust/sglang-radix-tree/src/python_bindings.rs index abfd80c48..9a81f2dc1 100644 --- a/rust/sglang-radix-tree/src/python_bindings.rs +++ b/rust/sglang-radix-tree/src/python_bindings.rs @@ -1775,6 +1775,17 @@ impl TreeCoreBinding { py.allow_threads(|| self.core().enable_storage) } + /// Enable or disable the direct external-cache linker. + fn set_enable_external_cache_linker(&self, py: Python<'_>, value: bool) -> PyResult<()> { + py.allow_threads(|| self.core().set_enable_external_cache_linker(value)) + .map_err(tree_core_assertion_error) + } + + /// Whether the direct external-cache linker is wired. + fn enable_external_cache_linker(&self, py: Python<'_>) -> bool { + py.allow_threads(|| self.core().enable_external_cache_linker) + } + /// Queue the all-cleared placement event. fn record_all_cleared_event(&self, py: Python<'_>) { py.allow_threads(|| self.core().record_all_cleared_event()); @@ -1872,6 +1883,64 @@ impl TreeCoreBinding { .map_err(node_access_error) } + /// Build transfers for a node with no stored or pending external copy. + fn build_external_linker_offload_transfers( + &self, + py: Python<'_>, + node_id: NodeId, + ) -> PyResult>>> { + let transfers = py + .allow_threads(|| self.core().build_external_linker_offload_transfers(node_id)) + .map_err(node_access_error)?; + transfers + .map(|transfers| { + transfers + .into_iter() + .map(|transfer| transfer_to_py(py, transfer)) + .collect::>>() + }) + .transpose() + } + + /// Mark an externally restored path, excluding its existing anchor. + fn mark_external_cache_stored_path( + &self, + py: Python<'_>, + from_node_id: NodeId, + until_node_id: NodeId, + ) -> PyResult<()> { + py.allow_threads(|| { + self.core() + .mark_external_cache_stored_path(from_node_id, until_node_id) + }) + .map_err(tree_core_runtime_error) + } + + /// Publish an accepted external offload as pending. + fn mark_external_linker_offload_pending( + &self, + py: Python<'_>, + node_id: NodeId, + ) -> PyResult<()> { + py.allow_threads(|| self.core().mark_external_linker_offload_pending(node_id)) + .map_err(tree_core_assertion_error) + } + + /// Finalize external-store state for an offload and its split fragments. + fn finish_external_linker_offload( + &self, + py: Python<'_>, + node_ids: Vec, + ack_id: NodeId, + success: bool, + ) -> PyResult<()> { + py.allow_threads(|| { + self.core() + .finish_external_linker_offload(&node_ids, ack_id, success) + }) + .map_err(tree_core_assertion_error) + } + /// Order-sensitive digest of reclaimed coexisting host values. fn write_back_coexist_reclaim_digest(&self, py: Python<'_>) -> i64 { py.allow_threads(|| self.core().write_back_coexist_reclaim_digest) @@ -1968,6 +2037,11 @@ impl TreeCoreBinding { .map_err(node_access_error) } + fn inspect_is_external_cache_stored(&self, py: Python<'_>, node_id: NodeId) -> PyResult { + py.allow_threads(|| self.core().inspect_is_external_cache_stored(node_id)) + .map_err(node_access_error) + } + fn inspect_is_node_in_device_lru( &self, py: Python<'_>, @@ -2839,6 +2913,20 @@ macro_rules! tree_core_binding { self.inner.enable_storage(py) } + /// Enable or disable the direct external-cache linker. + fn set_enable_external_cache_linker( + &self, + py: Python<'_>, + value: bool, + ) -> PyResult<()> { + self.inner.set_enable_external_cache_linker(py, value) + } + + /// Whether the direct external-cache linker is wired. + fn enable_external_cache_linker(&self, py: Python<'_>) -> bool { + self.inner.enable_external_cache_linker(py) + } + /// Queue the all-cleared placement event. fn record_all_cleared_event(&self, py: Python<'_>) { self.inner.record_all_cleared_event(py) @@ -2884,6 +2972,53 @@ macro_rules! tree_core_binding { self.inner.finish_load_back(py, anchor_node_id) } + /// Build transfers for a node with no stored or pending external copy. + fn build_external_linker_offload_transfers( + &self, + py: Python<'_>, + node_id: NodeId, + ) -> PyResult>>> { + self.inner + .build_external_linker_offload_transfers(py, node_id) + } + + /// Mark an externally restored path, excluding its existing anchor. + fn mark_external_cache_stored_path( + &self, + py: Python<'_>, + from_node_id: NodeId, + until_node_id: NodeId, + ) -> PyResult<()> { + self.inner.mark_external_cache_stored_path( + py, + from_node_id, + until_node_id, + ) + } + + /// Publish an accepted external offload as pending. + fn mark_external_linker_offload_pending( + &self, + py: Python<'_>, + node_id: NodeId, + ) -> PyResult<()> { + self.inner + .mark_external_linker_offload_pending(py, node_id) + } + + /// Finalize external-store state for an offload and its split fragments. + fn finish_external_linker_offload( + &self, + py: Python<'_>, + node_ids: Vec, + ack_id: NodeId, + success: bool, + ) -> PyResult<()> { + self.inner.finish_external_linker_offload( + py, node_ids, ack_id, success, + ) + } + /// Order-sensitive digest of reclaimed coexisting host values. #[pyo3(name = "write_back_duplicate_reclaim_digest")] fn write_back_coexist_reclaim_digest(&self, py: Python<'_>) -> i64 { @@ -2994,6 +3129,15 @@ macro_rules! tree_core_binding { .inspect_get_write_through_pending_id(py, node_id) } + #[cfg(feature = "inspection")] + fn inspect_is_external_cache_stored( + &self, + py: Python<'_>, + node_id: NodeId, + ) -> PyResult { + self.inner.inspect_is_external_cache_stored(py, node_id) + } + #[cfg(feature = "inspection")] fn inspect_is_node_in_device_lru( &self, diff --git a/rust/sglang-radix-tree/src/tests/unified_tree_core.rs b/rust/sglang-radix-tree/src/tests/unified_tree_core.rs index eb8644e5a..7ebc290a4 100644 --- a/rust/sglang-radix-tree/src/tests/unified_tree_core.rs +++ b/rust/sglang-radix-tree/src/tests/unified_tree_core.rs @@ -2201,6 +2201,248 @@ fn insert_threshold_crossing_emits_the_backup_kv_action() { assert_eq!(backups, vec![vec![leaf]]); } +#[test] +fn external_linker_hashes_new_nodes_and_triggers_offload_action() { + let params = CacheInitParams { + page_size: 2, + write_through_threshold: 1, + ..Default::default() + }; + let mut tc: UnifiedTreeCore> = UnifiedTreeCore::new(params, vec![FULL]); + tc.set_enable_external_cache_linker(true).unwrap(); + + let result = tc.insert(&insert_params(&vec![1, 2], &[10, 11])); + let leaf = result.last_device_node_id.unwrap(); + assert!(result.cache_actions.iter().any(|action| { + matches!(action, CacheAction::BackupKV(backup) if backup.node_ids == vec![leaf]) + })); + + let transfers = tc + .build_external_linker_offload_transfers(leaf) + .unwrap() + .unwrap(); + assert_eq!(transfers.len(), 1); + assert_eq!(transfers[0].name, PoolName::Kv); + assert_eq!(transfers[0].hit_policy, PoolHitPolicy::AllPages); + assert!( + transfers[0] + .device_indices + .as_ref() + .unwrap() + .equal(&Tensor::from_slice(&[10i64, 11])) + ); + assert_eq!( + transfers[0].keys, + tc.arena + .node(tc.arena.resolve(leaf).expect("live test node")) + .hash_value + .clone() + ); +} + +#[test] +fn external_linker_swa_offload_uses_complete_trailing_pages() { + let params = CacheInitParams { + page_size: 2, + swa_sliding_window_size: Some(4), + ..Default::default() + }; + let mut tc: UnifiedTreeCore> = UnifiedTreeCore::new(params, vec![FULL, SWA]); + tc.set_enable_external_cache_linker(true).unwrap(); + let root = tc.arena.root(); + let node = tc.add_new_node_( + root, + vec![1, 2, 3, 4, 5, 6], + &Tensor::from_slice(&[10i64, 11, 12, 13, 14, 15]), + 0, + None, + ); + tc.arena.node_mut(node).values[SWA.idx()].value = + Some(Tensor::from_slice(&[20i64, 21, 22, 23, 24])); + + let transfers = tc + .build_external_linker_offload_transfers(tc.arena.node(node).id) + .unwrap() + .unwrap(); + assert_eq!(transfers.len(), 2); + assert_eq!(transfers[0].name, PoolName::Kv); + assert_eq!(transfers[1].name, PoolName::Swa); + assert_eq!(transfers[1].hit_policy, PoolHitPolicy::TrailingPages); + assert!( + transfers[1] + .device_indices + .as_ref() + .unwrap() + .equal(&Tensor::from_slice(&[21i64, 22, 23, 24])) + ); + assert_eq!( + transfers[1].keys.as_ref().unwrap(), + &tc.arena.node(node).hash_value.as_ref().unwrap()[1..] + ); +} + +#[test] +fn external_linker_rejects_mamba_trees() { + let params = CacheInitParams { + mamba_cache_chunk_size: Some(1), + ..Default::default() + }; + let mut tc: UnifiedTreeCore> = UnifiedTreeCore::new(params, vec![FULL, MAMBA]); + let error = tc.set_enable_external_cache_linker(true).unwrap_err(); + assert!(matches!( + &error, + TreeCoreRuntimeError::ExternalCacheLinkerUnsupportedComponent { + component_type + } if *component_type == MAMBA + )); + assert!(error.to_string().contains("Mamba")); + assert!(!tc.enable_external_cache_linker); +} + +#[test] +fn external_linker_state_follows_load_offload_and_split_lifecycle() { + let mut tc = core(); + tc.set_enable_external_cache_linker(true).unwrap(); + tc.insert(&insert_params(&vec![1, 2], &[10, 11])); + tc.insert(&insert_params(&vec![1, 2, 3, 4], &[10, 11, 12, 13])); + let anchor = tc + .match_prefix(&match_params(&vec![1, 2])) + .best_match_node_id; + let leaf = tc + .match_prefix(&match_params(&vec![1, 2, 3, 4])) + .best_match_node_id; + + tc.mark_external_cache_stored_path(leaf, anchor).unwrap(); + assert!( + tc.arena + .node(tc.arena.resolve(leaf).expect("live test node")) + .external_cache_stored + ); + assert!( + !tc.arena + .node(tc.arena.resolve(anchor).expect("live test node")) + .external_cache_stored + ); + assert!( + tc.build_external_linker_offload_transfers(leaf) + .unwrap() + .is_none() + ); + assert!(matches!( + tc.mark_external_linker_offload_pending(leaf), + Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState { + node_id, + stored: true, + pending_id: None, + }) if node_id == leaf + )); + + let leaf_idx = tc.arena.resolve(leaf).expect("live test node"); + tc.arena.node_mut(leaf_idx).external_cache_stored = false; + tc.mark_external_linker_offload_pending(leaf).unwrap(); + assert!( + tc.build_external_linker_offload_transfers(leaf) + .unwrap() + .is_none() + ); + assert!(matches!( + tc.mark_external_linker_offload_pending(leaf), + Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState { + node_id, + stored: false, + pending_id: Some(pending_id), + }) if node_id == leaf && pending_id == leaf + )); + let (new_parent, action) = tc.split_node_(leaf_idx, 1); + let new_parent_handle = tc.arena.node(new_parent).id; + assert!(action.is_some()); + assert!(!tc.arena.node(new_parent).external_cache_stored); + assert!(!tc.arena.node(leaf_idx).external_cache_stored); + + tc.insert(&insert_params(&vec![9], &[19])); + let independent = tc.match_prefix(&match_params(&vec![9])).best_match_node_id; + tc.mark_external_linker_offload_pending(independent) + .unwrap(); + assert!(matches!( + tc.finish_external_linker_offload(&[independent, new_parent_handle], independent, false), + Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState { + node_id, + stored: false, + pending_id: Some(pending_id), + }) if node_id == new_parent_handle && pending_id == leaf + )); + assert_eq!( + tc.arena + .node(tc.arena.resolve(independent).expect("live test node")) + .write_through_pending_id, + Some(independent) + ); + for node_id in [new_parent_handle, leaf] { + let node = tc + .arena + .node(tc.arena.resolve(node_id).expect("live test node")); + assert_eq!(node.write_through_pending_id, Some(leaf)); + assert!(!node.external_cache_stored); + } + + tc.finish_external_linker_offload(&[independent], independent, false) + .unwrap(); + tc.finish_external_linker_offload(&[new_parent_handle, leaf], leaf, false) + .unwrap(); + for node_id in [new_parent_handle, leaf] { + let node = tc + .arena + .node(tc.arena.resolve(node_id).expect("live test node")); + assert_eq!(node.write_through_pending_id, None); + assert!(!node.external_cache_stored); + } +} + +#[test] +fn failed_external_offload_preserves_independently_confirmed_state() { + let mut tc = core(); + tc.set_enable_external_cache_linker(true).unwrap(); + tc.insert(&insert_params(&vec![1], &[10])); + tc.insert(&insert_params(&vec![1, 2], &[10, 11])); + let anchor = tc.match_prefix(&match_params(&vec![1])).best_match_node_id; + let leaf = tc + .match_prefix(&match_params(&vec![1, 2])) + .best_match_node_id; + + tc.mark_external_linker_offload_pending(leaf).unwrap(); + tc.mark_external_cache_stored_path(leaf, anchor).unwrap(); + tc.finish_external_linker_offload(&[leaf], leaf, false) + .unwrap(); + + let leaf = tc + .arena + .node(tc.arena.resolve(leaf).expect("live test node")); + assert_eq!(leaf.write_through_pending_id, None); + assert!(leaf.external_cache_stored); +} + +#[test] +fn external_linker_path_validation_is_atomic() { + let mut tc = core(); + tc.insert(&insert_params(&vec![1], &[10])); + tc.insert(&insert_params(&vec![2], &[20])); + let left = tc.match_prefix(&match_params(&vec![1])).best_match_node_id; + let right = tc.match_prefix(&match_params(&vec![2])).best_match_node_id; + + assert!(matches!( + tc.mark_external_cache_stored_path(left, right), + Err(TreeCoreRuntimeError::ExternalCachePathNotAncestor { + from_node_id, + until_node_id, + }) if from_node_id == left && until_node_id == right + )); + assert!( + !tc.arena + .node(tc.arena.resolve(left).expect("live test node")) + .external_cache_stored + ); +} + #[test] fn mark_write_through_pending_stamps_the_supplied_ack() { let mut tc = core(); @@ -2325,6 +2567,38 @@ fn backup_kv_action_chains_unbacked_ancestors_first() { assert_eq!(action.node_ids, vec![c]); } +#[test] +fn backup_kv_action_stops_at_an_externally_stored_or_pending_ancestor() { + let mut tc = core(); + tc.set_enable_external_cache_linker(true).unwrap(); + tc.insert(&insert_params(&vec![1], &[10])); + tc.insert(&insert_params(&vec![1, 2], &[10, 11])); + tc.insert(&insert_params(&vec![1, 2, 3], &[10, 11, 12])); + let a = tc.match_prefix(&match_params(&vec![1])).best_match_node_id; + let b = tc + .match_prefix(&match_params(&vec![1, 2])) + .best_match_node_id; + let c = tc + .match_prefix(&match_params(&vec![1, 2, 3])) + .best_match_node_id; + let a_idx = tc.arena.resolve(a).expect("live test node"); + tc.arena.node_mut(a_idx).external_cache_stored = true; + + let action = tc.build_backup_kv_action_( + tc.arena.node(tc.arena.resolve(c).expect("live test node")), + /* write_back = */ false, + ); + assert_eq!(action.node_ids, vec![b, c]); + + tc.arena.node_mut(a_idx).external_cache_stored = false; + tc.mark_external_linker_offload_pending(a).unwrap(); + let action = tc.build_backup_kv_action_( + tc.arena.node(tc.arena.resolve(c).expect("live test node")), + /* write_back = */ false, + ); + assert_eq!(action.node_ids, vec![b, c]); +} + #[test] fn split_of_a_pending_node_transfers_the_ack_and_emits_the_replace_action() { let mut tc = core(); @@ -3958,6 +4232,35 @@ fn fallible_node_boundaries_reject_stale_handles() { Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id })) if node_id == stale_root )); + assert!(matches!( + tc.build_external_linker_offload_transfers(stale_root), + Err(NodeAccessError { node_id }) if node_id == stale_root + )); + assert!(matches!( + tc.mark_external_cache_stored_path(stale_root, live_root), + Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id })) + if node_id == stale_root + )); + assert!(matches!( + tc.mark_external_cache_stored_path(live_root, stale_root), + Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id })) + if node_id == stale_root + )); + assert!(matches!( + tc.mark_external_linker_offload_pending(stale_root), + Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id })) + if node_id == stale_root + )); + assert!(matches!( + tc.finish_external_linker_offload(&[live_root, stale_root], live_root, true), + Err(TreeCoreRuntimeError::NodeAccess(NodeAccessError { node_id })) + if node_id == stale_root + )); + assert!( + !tc.arena + .node(tc.arena.resolve(live_root).expect("live root")) + .external_cache_stored + ); assert!(matches!( tc.get_hash_values(stale_root), Err(NodeAccessError { node_id }) if node_id == stale_root @@ -7964,6 +8267,10 @@ fn inspection_rejects_stale_handles_without_panicking() { assert_eq!(tc.inspect_get_parent_node_id(stale_root), Err(expected)); assert_eq!(tc.inspect_get_child_node_ids(stale_root), Err(expected)); assert_eq!(tc.inspect_get_node_key_length(stale_root), Err(expected)); + assert_eq!( + tc.inspect_is_external_cache_stored(stale_root), + Err(expected) + ); assert_eq!( tc.inspect_set_node_hash_values(stale_root, None), Err(expected) @@ -7993,6 +8300,7 @@ fn inspection_rejects_stale_handles_without_panicking() { assert!(!tc.inspect_is_host_evictable_leaf(stale_root)); assert_eq!(tc.inspect_get_parent_node_id(live_root), Ok(None)); + assert_eq!(tc.inspect_is_external_cache_stored(live_root), Ok(false)); } #[test] diff --git a/rust/sglang-radix-tree/src/unified_tree_core.rs b/rust/sglang-radix-tree/src/unified_tree_core.rs index 33530f60c..c19713914 100644 --- a/rust/sglang-radix-tree/src/unified_tree_core.rs +++ b/rust/sglang-radix-tree/src/unified_tree_core.rs @@ -534,6 +534,8 @@ pub struct UnifiedTreeCore { pub(crate) enable_hicache: bool, /// Whether the storage tier (L3) is wired; gates page-hash computation. pub(crate) enable_storage: bool, + /// Whether a direct device-to-external-cache linker is wired. + pub(crate) enable_external_cache_linker: bool, /// Whether the cache wired a host SWA pool (HiCache). pub(crate) has_swa_host_pool: bool, /// Whether tree mutations emit BlockStored/BlockRemoved events. @@ -717,6 +719,7 @@ impl UnifiedTreeCore { is_write_back: params.is_write_back, enable_hicache: params.enable_hicache, enable_storage: false, + enable_external_cache_linker: false, has_swa_host_pool: params.has_swa_host_pool, enable_kv_cache_events: params.enable_kv_cache_events, kv_event_queue: Vec::new(), @@ -1319,6 +1322,12 @@ impl UnifiedTreeCore { return false; } node.hit_count += 1; + + if self.enable_external_cache_linker { + return Self::needs_external_linker_offload_(node) + && node.hit_count >= self.write_through_threshold; + } + self.enable_hicache && !node.backuped() && node.hit_count >= self.write_through_threshold } @@ -1763,6 +1772,7 @@ impl UnifiedTreeCore { let child = self.arena.node(child_id); let parent_id = child.parent(); let child_namespace = child.namespace.clone(); + let child_external_cache_stored = child.external_cache_stored; let (key_head, key_tail) = child.key.split_at(split_len); // key_head keeps the original key's first page, which keys the parent's child map. let parent_map_key = key_head.child_key(page_size); @@ -1778,6 +1788,7 @@ impl UnifiedTreeCore { (child_namespace.clone(), key_tail.child_key(page_size)), child_id, ); + self.arena.node_mut(new_node_id).external_cache_stored = child_external_cache_stored; // The child's aux LRU cells detach while it is re-linked. self.for_each_component_lru_( @@ -1897,7 +1908,7 @@ impl UnifiedTreeCore { "add_new_node_: parent {parent_id} already has a child on the new node's page" ); self.inc_evictable_size(FULL, value.size()[0] as usize); - if self.enable_storage { + if self.enable_storage || self.enable_external_cache_linker { let hash_values = self.arena.compute_node_hash_values(new_node_id, page_size); self.arena.node_mut(new_node_id).hash_value = Some(hash_values); } @@ -2708,6 +2719,22 @@ impl UnifiedTreeCore { self.enable_storage = value; } + /// Enable or disable the direct external-cache linker. + pub fn set_enable_external_cache_linker( + &mut self, + value: bool, + ) -> Result<(), TreeCoreRuntimeError> { + if value && self.components_by_type[MAMBA.idx()].is_some() { + return Err( + TreeCoreRuntimeError::ExternalCacheLinkerUnsupportedComponent { + component_type: MAMBA, + }, + ); + } + self.enable_external_cache_linker = value; + Ok(()) + } + // ==== KV cache placement events ==== /// Append an event, coalescing it with a compatible queue tail. @@ -3463,14 +3490,19 @@ impl UnifiedTreeCore { Ok(order) } - /// Build the backup action for a node and its unbacked ancestors. + /// Build the backup action for a node and its not-yet-persisted ancestors. pub fn build_backup_kv_action_(&self, node: &Node, write_back: bool) -> BackupKV { let mut chain = vec![node.id]; if !write_back { let mut ancestor = node.try_parent(); while let Some(ancestor_idx) = ancestor { let ancestor_node = self.arena.node(ancestor_idx); - if ancestor_node.is_root() || ancestor_node.backuped() { + if ancestor_node.is_root() + || ancestor_node.backuped() + || ancestor_node.external_cache_stored + || (self.enable_external_cache_linker + && ancestor_node.write_through_pending_id.is_some()) + { break; } chain.push(ancestor_node.id); @@ -3696,6 +3728,103 @@ impl UnifiedTreeCore { depth } + /// Build transfers for a node with no stored or pending external copy. + pub fn build_external_linker_offload_transfers( + &self, + node_id: NodeId, + ) -> Result>, NodeAccessError> { + let node_id = self.arena.resolve(node_id)?; + if !Self::needs_external_linker_offload_(self.arena.node(node_id)) { + return Ok(None); + } + + let transfers = self + .components + .iter() + .filter_map(|component| component.build_external_linker_offload_transfer(self, node_id)) + .collect(); + Ok(Some(transfers)) + } + + fn needs_external_linker_offload_(node: &Node) -> bool { + !node.external_cache_stored && node.write_through_pending_id.is_none() + } + + /// Mark the path from `from_node_id` to, but excluding, `until_node_id` as + /// available in the external cache. + pub fn mark_external_cache_stored_path( + &mut self, + from_node_id: NodeId, + until_node_id: NodeId, + ) -> Result<(), TreeCoreRuntimeError> { + let from = self.arena.resolve(from_node_id)?; + let until = self.arena.resolve(until_node_id)?; + let mut path = Vec::new(); + let mut current = from; + while current != until { + let node = self.arena.node(current); + let Some(parent) = node.try_parent() else { + return Err(TreeCoreRuntimeError::ExternalCachePathNotAncestor { + from_node_id, + until_node_id, + }); + }; + path.push(current); + current = parent; + } + for node_id in path { + self.arena.node_mut(node_id).external_cache_stored = true; + } + Ok(()) + } + + /// Publish an accepted external offload as pending. + pub fn mark_external_linker_offload_pending( + &mut self, + node_id: NodeId, + ) -> Result<(), TreeCoreRuntimeError> { + let node_idx = self.arena.resolve(node_id)?; + let node = self.arena.node(node_idx); + if !Self::needs_external_linker_offload_(node) { + return Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState { + node_id, + stored: node.external_cache_stored, + pending_id: node.write_through_pending_id, + }); + } + self.arena.node_mut(node_idx).write_through_pending_id = Some(node_id); + Ok(()) + } + + /// Finalize external-store state for an offload and its split fragments. + pub fn finish_external_linker_offload( + &mut self, + node_ids: &[NodeId], + ack_id: NodeId, + success: bool, + ) -> Result<(), TreeCoreRuntimeError> { + let node_indices = node_ids + .iter() + .map(|&node_id| self.arena.resolve(node_id)) + .collect::, _>>()?; + for (&node_id, &node_idx) in node_ids.iter().zip(&node_indices) { + let node = self.arena.node(node_idx); + if node.write_through_pending_id != Some(ack_id) { + return Err(TreeCoreRuntimeError::InvalidExternalCacheOffloadState { + node_id, + stored: node.external_cache_stored, + pending_id: node.write_through_pending_id, + }); + } + } + for node_id in node_indices { + let node = self.arena.node_mut(node_id); + node.write_through_pending_id = None; + node.external_cache_stored |= success; + } + Ok(()) + } + /// Clear the write-through-pending mark (when it matches ack_id) and record the /// host store event for each acked node. pub fn finish_write_through( @@ -4357,6 +4486,15 @@ impl UnifiedTreeCore { Ok(self.arena.node(node_id).write_through_pending_id) } + /// Whether a node is known to be stored in the external cache. + pub fn inspect_is_external_cache_stored( + &self, + node_id: NodeId, + ) -> Result { + let node_id = self.arena.resolve(node_id)?; + Ok(self.arena.node(node_id).external_cache_stored) + } + /// Whether a node is in a component's device LRU. pub fn inspect_is_node_in_device_lru( &self, diff --git a/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_glm52.py b/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_glm52.py index 6f56b85f2..50792565d 100644 --- a/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_glm52.py +++ b/test/registered/radix_cache/unified_radix_tree/linker/test_unified_cache_linker_kl_glm52.py @@ -13,6 +13,7 @@ from sglang.test.test_utils import ( find_available_port, popen_launch_server, terminate_and_kill_process_tree, + unified_radix_tree_server_env, ) GLM52_MODEL = os.environ.get("SGLANG_LINKER_GLM52_MODEL", "zai-org/GLM-5.2-FP8") @@ -22,6 +23,7 @@ register_cuda_ci(est_time=383, stage="extra-b", runner_config="8-gpu-h200") class TestGLM52UnifiedCacheLinkerKL(UnifiedRadixTreeTestMixin, CustomTestCase): + tree_core_backend = "rust" page_size = 64 kl_threshold = 0.03 sampling_temperature = 0 @@ -61,10 +63,10 @@ class TestGLM52UnifiedCacheLinkerKL(UnifiedRadixTreeTestMixin, CustomTestCase): "--hicache-storage-backend-extra-config", json.dumps({"enable_group_semantics": True}), ], - env={ + env=unified_radix_tree_server_env( + cls.tree_core_backend, **cls.mooncake.server_env(), - "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", - }, + ), ) cls.input_ids = get_input_ids(cls.model, num_samples=18) except Exception: diff --git a/test/registered/unit/mem_cache/rust_unified_tree_core_inspector.py b/test/registered/unit/mem_cache/rust_unified_tree_core_inspector.py index 097e4ffa9..5b4d7daa5 100644 --- a/test/registered/unit/mem_cache/rust_unified_tree_core_inspector.py +++ b/test/registered/unit/mem_cache/rust_unified_tree_core_inspector.py @@ -82,6 +82,9 @@ class RustUnifiedTreeCoreInspector( def get_write_through_pending_id(self, node_id: NodeId) -> Optional[int]: return self._binding.inspect_get_write_through_pending_id(node_id) + def is_external_cache_stored(self, node_id: NodeId) -> bool: + return self._binding.inspect_is_external_cache_stored(node_id) + def is_node_in_device_lru( self, node_id: NodeId, component_type: ComponentType ) -> bool: diff --git a/test/registered/unit/mem_cache/test_rust_tree_core_integration.py b/test/registered/unit/mem_cache/test_rust_tree_core_integration.py index f376dfdb0..9c935379f 100644 --- a/test/registered/unit/mem_cache/test_rust_tree_core_integration.py +++ b/test/registered/unit/mem_cache/test_rust_tree_core_integration.py @@ -266,6 +266,21 @@ def test_stale_handle_operations_raise_key_error_without_poisoning_the_core(): stale_root, empty, PoolTransfer(name=PoolName.KV), {} ), "build_load_back_spec": lambda: core.build_load_back_spec(stale_root), + "build_external_linker_offload_transfers": lambda: ( + core.build_external_linker_offload_transfers(stale_root) + ), + "mark_external_cache_stored_path/from": lambda: ( + core.mark_external_cache_stored_path(stale_root, live_root) + ), + "mark_external_cache_stored_path/until": lambda: ( + core.mark_external_cache_stored_path(live_root, stale_root) + ), + "mark_external_linker_offload_pending": lambda: ( + core.mark_external_linker_offload_pending(stale_root) + ), + "finish_external_linker_offload": lambda: core.finish_external_linker_offload( + [live_root, stale_root], live_root, True + ), "evict_excess_path_states": lambda: core.evict_excess_path_states( stale_root, {}, {} ), @@ -499,13 +514,19 @@ def test_configuration_reads_the_locked_rust_state(): assert swa_core.has_swa_host_pool is True -def test_external_cache_linker_is_rejected(): +def test_external_cache_linker_enablement_and_component_guard(): core = _tree_core() assert core.enable_external_cache_linker is False - with pytest.raises(ValueError, match="External cache linker"): - core.enable_external_cache_linker = True + core.enable_external_cache_linker = True + assert core.enable_external_cache_linker is True + core.enable_external_cache_linker = False assert core.enable_external_cache_linker is False + mamba_core = _mamba_tree_core() + with pytest.raises(AssertionError, match="(?i)mamba"): + mamba_core.enable_external_cache_linker = True + assert mamba_core.enable_external_cache_linker is False + def test_sanity_check_passes_after_the_full_flow(): core = _tree_core() @@ -2179,6 +2200,7 @@ def test_stale_inspection_handles_raise_key_error_or_report_absence(): "get_write_through_pending_id": lambda: core.get_write_through_pending_id( stale_root ), + "is_external_cache_stored": lambda: core.is_external_cache_stored(stale_root), "is_node_in_device_lru": lambda: core.is_node_in_device_lru( stale_root, ComponentType.FULL ), diff --git a/test/registered/unit/mem_cache/test_unified_cache_linker.py b/test/registered/unit/mem_cache/test_unified_cache_linker.py index 7082619dc..fa48f6638 100644 --- a/test/registered/unit/mem_cache/test_unified_cache_linker.py +++ b/test/registered/unit/mem_cache/test_unified_cache_linker.py @@ -1,10 +1,33 @@ +import sys +import unittest +from array import array +from collections import defaultdict +from dataclasses import replace from types import SimpleNamespace import pytest +import test_unified_radix_cache_unittest as shared_cache_suite import torch +from test_unified_radix_cache_unittest import ( + CacheConfig, + _device_lock_ref, + _device_value, + _InsertWalkSuite, + _node_children, + build_fixture, +) -from sglang.srt.mem_cache.base_prefix_cache import InsertResult -from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer +from sglang.srt.mem_cache.base_prefix_cache import ( + InitLoadBackParams, + InsertResult, + MatchPrefixParams, +) +from sglang.srt.mem_cache.hicache_storage import ( + PoolHitPolicy, + PoolName, + PoolTransfer, +) +from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.unified_cache.cache_action import ( ReplaceWriteThroughOnNodeSplit, ) @@ -13,16 +36,16 @@ from sglang.srt.mem_cache.unified_cache.components.full_component import FullCom from sglang.srt.mem_cache.unified_cache.components.swa_component import SWAComponent from sglang.srt.mem_cache.unified_cache.components.tree_component import ( ExternalLinkerLoadPhase, - LinkerTransferPhase, ) from sglang.srt.mem_cache.unified_cache.unified_cache_linker import ( UnifiedCacheLinker, UnifiedCacheLinkerWrapper, ) from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache -from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.ci.ci_register import register_cpu_ci, register_cuda_ci register_cpu_ci(est_time=10, suite="base-a-test-cpu") +register_cuda_ci(est_time=100, stage="base-b", runner_config="1-gpu-small") class _FakeLinker(UnifiedCacheLinker): @@ -83,9 +106,41 @@ class _MappingRecorder: self.mapping.append((full.clone(), swa.clone())) +class _FakeExternalTreeCore: + def __init__(self, nodes=None, offload_transfers=None): + self.enable_external_cache_linker = False + self.nodes = nodes or {} + self.offload_transfers = offload_transfers or [ + PoolTransfer(name=PoolName.KV, keys=["page"]) + ] + + def build_external_linker_offload_transfers(self, node_id): + node = self.nodes[node_id] + if node.external_cache_stored or node.write_through_pending_id is not None: + return None + return list(self.offload_transfers) + + def mark_external_linker_offload_pending(self, node_id): + node = self.nodes[node_id] + assert ( + not node.external_cache_stored and node.write_through_pending_id is None + ), "invalid external offload state" + node.write_through_pending_id = node_id + + def finish_external_linker_offload(self, node_ids, ack_id, success): + nodes = [self.nodes[node_id] for node_id in node_ids] + assert all(node.write_through_pending_id == ack_id for node in nodes), ( + "invalid external offload state" + ) + for node in nodes: + node.write_through_pending_id = None + node.external_cache_stored |= success + + def _cache_for_wrapper(**kwargs): defaults = { "tree_core": SimpleNamespace(enable_external_cache_linker=False), + "tree_components": (ComponentType.FULL,), "write_through_threshold": 256, "pp_size": 1, "pp_group": None, @@ -100,6 +155,7 @@ def test_cache_linker_attachment_is_backend_independent(): enable_external_cache_linker=False, write_through_threshold=256, ) + cache.tree_components = (ComponentType.FULL,) cache.linker = None linker = _FakeLinker() @@ -111,6 +167,585 @@ def test_cache_linker_attachment_is_backend_independent(): assert cache.linker.layer_done_counter is linker.layer_done_counter +@pytest.mark.parametrize("component_type", [ComponentType.MAMBA, ComponentType.C128]) +def test_cache_linker_rejects_unsupported_tree_components(component_type): + cache = _cache_for_wrapper(tree_components=(ComponentType.FULL, component_type)) + + with pytest.raises(ValueError, match=component_type.name): + UnifiedCacheLinkerWrapper(cache, _FakeLinker()) + + assert not cache.tree_core.enable_external_cache_linker + + +class _InMemoryUnifiedCacheLinker(UnifiedCacheLinker): + """Controllable transport for shared Python/Rust TreeCore tests.""" + + def __init__(self, stored_keys=None): + self.layer_done_counter = object() + self.stored_keys = defaultdict(set) if stored_keys is None else stored_keys + self.lookup_calls = [] + self.offload_calls = [] + self.pending_offloads = [] + self.queued_loads = {} + self.started_loads = [] + self.completed_loads = [] + self.completed_offloads = [] + + @staticmethod + def _clone_transfer(transfer): + return replace( + transfer, + host_indices=( + None if transfer.host_indices is None else transfer.host_indices.clone() + ), + device_indices=( + None + if transfer.device_indices is None + else transfer.device_indices.clone() + ), + keys=None if transfer.keys is None else list(transfer.keys), + nodes_to_load=( + None if transfer.nodes_to_load is None else list(transfer.nodes_to_load) + ), + ) + + @classmethod + def _clone_transfers(cls, transfers): + return [cls._clone_transfer(transfer) for transfer in transfers] + + def lookup(self, rid, transfers): + transfers = self._clone_transfers(transfers) + self.lookup_calls.append((rid, transfers)) + by_pool = {transfer.name: transfer for transfer in transfers} + kv = by_pool.get(PoolName.KV) + if kv is None or not kv.keys: + return [] + + restorable = [] + for prefix_pages in range(1, len(kv.keys) + 1): + if not set(kv.keys[:prefix_pages]) <= self.stored_keys[PoolName.KV]: + continue + for transfer in transfers: + if transfer.name == PoolName.KV: + continue + if transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES: + window_pages = max(1, len(transfer.keys or ())) + required = kv.keys[ + max(0, prefix_pages - window_pages) : prefix_pages + ] + else: + required = kv.keys[:prefix_pages] + if not set(required) <= self.stored_keys[transfer.name]: + break + else: + restorable.append(prefix_pages) + return restorable + + def load(self, rid, transfers): + self.queued_loads[rid] = self._clone_transfers(transfers) + return True + + def start_layer_wise_loading(self): + if not self.queued_loads: + return -1 + rids = list(self.queued_loads) + self.started_loads.append(rids) + return len(self.started_loads) - 1 + + def cancel_queued_load(self, rid): + return self.queued_loads.pop(rid, None) is not None + + def num_completed_loads(self): + return len(self.completed_loads) + + def pop_completed_load(self): + rids = self.completed_loads.pop(0) + for rid in rids: + self.queued_loads.pop(rid, None) + return rids + + def complete_started_loads(self): + self.completed_loads.append(self.started_loads[-1]) + + def offload(self, transfers): + transfers = self._clone_transfers(transfers) + self.offload_calls.append(transfers) + self.pending_offloads.append(transfers) + return True + + def num_completed_offloads(self): + return len(self.completed_offloads) + + def pop_completed_offload(self): + return self.completed_offloads.pop(0) + + def complete_next_offload(self, success): + transfers = self.pending_offloads.pop(0) + if success: + for transfer in transfers: + self.stored_keys[transfer.name].update(transfer.keys or ()) + self.completed_offloads.append(success) + + def reset(self): + self.queued_loads.clear() + self.pending_offloads.clear() + self.completed_loads.clear() + self.completed_offloads.clear() + + def close(self): + self.reset() + + +class _TreeCoreBackendTestMixin: + tree_core_backend = "python" + + def setUp(self): + previous = shared_cache_suite._TREE_CORE_TEST_BACKEND + self.addCleanup( + setattr, + shared_cache_suite, + "_TREE_CORE_TEST_BACKEND", + previous, + ) + shared_cache_suite._TREE_CORE_TEST_BACKEND = self.tree_core_backend + super().setUp() + + +@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA") +class TestUnifiedCacheLinkerPythonBackend(_TreeCoreBackendTestMixin, _InsertWalkSuite): + def test_full_offload_load_round_trip_and_dedup(self): + cfg = CacheConfig(page_size=2, kv_size=64, max_context_len=64) + self.cfg = cfg + stored_keys = defaultdict(set) + + producer, producer_allocator, producer_req_pool = build_fixture(cfg) + producer_linker = _InMemoryUnifiedCacheLinker(stored_keys) + producer.init_cache_linker(producer_linker) + tokens = list(range(1, 9)) + + inserted = self._insert(producer, producer_allocator, producer_req_pool, tokens) + self.assertEqual(len(producer_linker.offload_calls), 1) + (kv_offload,) = producer_linker.offload_calls[0] + self.assertEqual(kv_offload.name, PoolName.KV) + self.assertEqual( + kv_offload.keys, + producer.tree_core.get_hash_values(inserted.last_device_node), + ) + self.assertTrue( + torch.equal( + kv_offload.device_indices, + _device_value(producer, inserted.last_device_node, ComponentType.FULL), + ) + ) + self.assertEqual( + _device_lock_ref(producer, inserted.last_device_node, ComponentType.FULL), + 1, + ) + self.assertFalse( + producer.tree_core.is_external_cache_stored(inserted.last_device_node) + ) + self.assertEqual( + producer.tree_core.get_write_through_pending_id(inserted.last_device_node), + inserted.last_device_node, + ) + + self._insert(producer, producer_allocator, producer_req_pool, tokens) + self.assertEqual(len(producer_linker.offload_calls), 1) + + producer_linker.complete_next_offload(True) + producer.check_hicache_events() + self.assertEqual( + _device_lock_ref(producer, inserted.last_device_node, ComponentType.FULL), + 0, + ) + self.assertTrue( + producer.tree_core.is_external_cache_stored(inserted.last_device_node) + ) + self._insert(producer, producer_allocator, producer_req_pool, tokens) + self.assertEqual(len(producer_linker.offload_calls), 1) + + consumer, _, consumer_req_pool = build_fixture(cfg) + consumer_linker = _InMemoryUnifiedCacheLinker(stored_keys) + consumer.init_cache_linker(consumer_linker) + req = self._make_req(consumer_req_pool) + match = consumer.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", tokens)), req=req) + ) + self.assertEqual(match.device_indices.numel(), 0) + self.assertEqual(match.host_hit_length, len(tokens)) + self._apply_match_to_req(req, match) + + loaded, loaded_node = consumer.init_load_back( + InitLoadBackParams( + best_match_node=match.best_match_node, + host_hit_length=match.host_hit_length, + req=req, + ) + ) + self.assertEqual(loaded.numel(), len(tokens)) + self.assertNotEqual(loaded_node, consumer.root_node_handle()) + (kv_load,) = consumer_linker.queued_loads[req.rid] + self.assertEqual(kv_load.name, PoolName.KV) + self.assertEqual(kv_load.keys, kv_offload.keys) + self.assertEqual(_device_lock_ref(consumer, loaded_node, ComponentType.FULL), 1) + + self.assertGreaterEqual(consumer.ready_to_load_host_cache(), 0) + consumer_linker.complete_started_loads() + consumer.check_hicache_events() + self.assertEqual(_device_lock_ref(consumer, loaded_node, ComponentType.FULL), 0) + final_match = consumer.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", tokens))) + ) + self.assertEqual(final_match.device_indices.numel(), len(tokens)) + self.assertEqual(consumer_linker.offload_calls, []) + consumer.sanity_check() + + def test_eagle_lookup_uses_bigram_tail_hashes(self): + cfg = CacheConfig( + page_size=2, + is_eagle=True, + kv_size=64, + max_context_len=64, + ) + self.cfg = cfg + stored_keys = defaultdict(set) + tokens = list(range(1, 10)) + + producer, producer_allocator, producer_req_pool = build_fixture(cfg) + producer_linker = _InMemoryUnifiedCacheLinker(stored_keys) + producer.init_cache_linker(producer_linker) + self._insert(producer, producer_allocator, producer_req_pool, tokens) + self.assertEqual(len(producer_linker.offload_calls), 1) + (kv_offload,) = producer_linker.offload_calls[0] + self.assertEqual(len(kv_offload.keys), 4) + producer_linker.complete_next_offload(True) + producer.check_hicache_events() + + consumer, consumer_allocator, consumer_req_pool = build_fixture(cfg) + consumer_linker = _InMemoryUnifiedCacheLinker(stored_keys) + consumer.init_cache_linker(consumer_linker) + consumer.write_through_threshold = sys.maxsize + self._insert( + consumer, + consumer_allocator, + consumer_req_pool, + tokens[:5], + ) + req = self._make_req(consumer_req_pool) + lookup_key = RadixKey(array("q", tokens)) + + match = consumer.match_prefix(MatchPrefixParams(key=lookup_key, req=req)) + + self.assertTrue(lookup_key.is_bigram) + self.assertEqual(match.device_indices.numel(), 4) + self.assertEqual(match.host_hit_length, 4) + self.assertEqual(len(consumer_linker.lookup_calls), 1) + _, transfers = consumer_linker.lookup_calls[0] + (kv_lookup,) = transfers + self.assertEqual(kv_lookup.keys, kv_offload.keys[2:]) + + def test_failed_split_offload_retries_and_reset_releases_locks(self): + cfg = CacheConfig(page_size=1, kv_size=64, max_context_len=64) + self.cfg = cfg + cache, allocator, req_to_token_pool = build_fixture(cfg) + linker = _InMemoryUnifiedCacheLinker() + cache.init_cache_linker(linker) + tokens = [1, 2, 3, 4] + + inserted = self._insert(cache, allocator, req_to_token_pool, tokens) + original_node = inserted.last_device_node + self.assertEqual(len(linker.offload_calls), 1) + original_keys = linker.offload_calls[0][0].keys + + self._insert(cache, allocator, req_to_token_pool, tokens[:2]) + (parent,) = _node_children(cache, cache.root_node_handle()) + (child,) = _node_children(cache, parent) + self.assertEqual(child, original_node) + self.assertEqual( + cache.linker.pending_offloads[0].publish_node_ids, [parent, child] + ) + self.assertEqual(cache.tree_core.get_write_through_pending_id(parent), child) + self.assertEqual(cache.tree_core.get_write_through_pending_id(child), child) + self.assertFalse(cache.tree_core.is_external_cache_stored(parent)) + self.assertFalse(cache.tree_core.is_external_cache_stored(child)) + + linker.complete_next_offload(False) + cache.check_hicache_events() + for node_id in (parent, child): + self.assertIsNone(cache.tree_core.get_write_through_pending_id(node_id)) + self.assertFalse(cache.tree_core.is_external_cache_stored(node_id)) + self.assertEqual(_device_lock_ref(cache, node_id, ComponentType.FULL), 0) + + self._insert(cache, allocator, req_to_token_pool, tokens) + retry_calls = linker.offload_calls[1:] + self.assertEqual(len(retry_calls), 2) + self.assertEqual( + [key for transfers in retry_calls for key in transfers[0].keys], + original_keys, + ) + for _ in retry_calls: + linker.complete_next_offload(True) + cache.check_hicache_events() + for node_id in (parent, child): + self.assertIsNone(cache.tree_core.get_write_through_pending_id(node_id)) + self.assertTrue(cache.tree_core.is_external_cache_stored(node_id)) + self.assertEqual(_device_lock_ref(cache, node_id, ComponentType.FULL), 0) + + self._insert(cache, allocator, req_to_token_pool, tokens) + self.assertEqual(len(linker.offload_calls), 3) + + extended = self._insert(cache, allocator, req_to_token_pool, tokens + [5, 6]) + pending_node = extended.last_device_node + self.assertEqual(len(cache.linker.pending_offloads), 1) + self.assertEqual( + cache.tree_core.get_write_through_pending_id(pending_node), pending_node + ) + self.assertEqual(_device_lock_ref(cache, pending_node, ComponentType.FULL), 1) + + cache.linker.reset() + self.assertEqual(cache.linker.pending_offloads, []) + self.assertIsNone(cache.tree_core.get_write_through_pending_id(pending_node)) + self.assertEqual(_device_lock_ref(cache, pending_node, ComponentType.FULL), 0) + cache.sanity_check() + cache.reset() + cache.sanity_check() + + def test_swa_partial_hit_loads_only_pages_not_adopted_locally(self): + cfg = CacheConfig( + page_size=1, + components=(ComponentType.FULL, ComponentType.SWA), + sliding_window_size=2, + kv_size=64, + max_context_len=64, + ) + self.cfg = cfg + stored_keys = defaultdict(set) + tokens = list(range(1, 7)) + + producer, producer_allocator, producer_req_pool = build_fixture(cfg) + producer_linker = _InMemoryUnifiedCacheLinker(stored_keys) + producer.init_cache_linker(producer_linker) + self._insert(producer, producer_allocator, producer_req_pool, tokens) + self.assertGreaterEqual(len(producer_linker.offload_calls), 1) + all_keys = [] + for transfers in producer_linker.offload_calls: + offloads = {transfer.name: transfer for transfer in transfers} + self.assertEqual(set(offloads), {PoolName.KV, PoolName.SWA}) + self.assertEqual(offloads[PoolName.KV].keys, offloads[PoolName.SWA].keys) + all_keys.extend(offloads[PoolName.KV].keys) + producer_linker.complete_next_offload(True) + producer.check_hicache_events() + + stored_keys[PoolName.KV].difference_update(all_keys[-2:]) + + consumer, consumer_allocator, consumer_req_pool = build_fixture(cfg) + consumer_linker = _InMemoryUnifiedCacheLinker(stored_keys) + consumer.init_cache_linker(consumer_linker) + consumer.write_through_threshold = sys.maxsize + self._insert(consumer, consumer_allocator, consumer_req_pool, tokens[:2]) + + req = self._make_req(consumer_req_pool) + match = consumer.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", tokens)), req=req) + ) + self.assertEqual(match.device_indices.numel(), 2) + self.assertEqual(match.host_hit_length, 2) + self.assertEqual(match.swa_host_hit_length, 2) + self._apply_match_to_req(req, match) + + self._insert(consumer, consumer_allocator, consumer_req_pool, tokens[:3]) + raced_match = consumer.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", tokens[:3]))) + ) + raced_full = raced_match.device_indices[-1:].clone() + raced_swa = consumer_allocator.translate_loc_from_full_to_swa(raced_full) + + loaded, loaded_node = consumer.init_load_back( + InitLoadBackParams( + best_match_node=match.best_match_node, + host_hit_length=match.host_hit_length, + req=req, + ) + ) + self.assertEqual(loaded.numel(), 2) + self.assertTrue(torch.equal(loaded[:1], raced_full)) + load_by_pool = { + transfer.name: transfer + for transfer in consumer_linker.queued_loads[req.rid] + } + self.assertEqual(set(load_by_pool), {PoolName.KV, PoolName.SWA}) + expected_key = all_keys[3] + for transfer in load_by_pool.values(): + self.assertEqual(transfer.keys, [expected_key]) + self.assertEqual(transfer.device_indices.numel(), 1) + + translated = consumer_allocator.translate_loc_from_full_to_swa(loaded) + self.assertTrue(torch.equal(translated[:1], raced_swa)) + self.assertTrue( + torch.equal(translated[-1:], load_by_pool[PoolName.SWA].device_indices) + ) + + self.assertGreaterEqual(consumer.ready_to_load_host_cache(), 0) + consumer_linker.complete_started_loads() + consumer.check_hicache_events() + final_match = consumer.match_prefix( + MatchPrefixParams(key=RadixKey(array("q", tokens[:4]))) + ) + self.assertEqual(final_match.device_indices.numel(), 4) + self.assertEqual(_device_lock_ref(consumer, loaded_node, ComponentType.FULL), 0) + consumer.sanity_check() + + +@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA") +class TestUnifiedCacheLinkerTreeCorePythonBackend( + _TreeCoreBackendTestMixin, shared_cache_suite._InsertWalkSuite +): + """TreeCore linker contracts shared by the Python and Rust inspectors.""" + + def test_builds_opaque_external_offload_transfers(self): + cfg = shared_cache_suite.CacheConfig( + page_size=2, kv_size=64, max_context_len=64 + ) + self.cfg = cfg + cache, allocator, req_to_token_pool = shared_cache_suite.build_fixture(cfg) + core = cache.tree_core + core.enable_external_cache_linker = True + + inserted = self._insert(cache, allocator, req_to_token_pool, list(range(1, 9))) + node_id = inserted.last_device_node + transfers = core.build_external_linker_offload_transfers(node_id) + + self.assertIsNotNone(transfers) + (transfer,) = transfers + self.assertEqual(transfer.name, PoolName.KV) + self.assertIsNone(transfer.host_indices) + self.assertIsNotNone(transfer.device_indices) + self.assertEqual(transfer.keys, core.get_hash_values(node_id)) + self.assertTrue( + torch.equal( + transfer.device_indices, + shared_cache_suite._device_value(cache, node_id, ComponentType.FULL), + ) + ) + + core.mark_external_linker_offload_pending(node_id) + self.assertFalse(core.is_external_cache_stored(node_id)) + self.assertIsNone(core.build_external_linker_offload_transfers(node_id)) + with self.assertRaisesRegex(AssertionError, "invalid external offload state"): + core.mark_external_linker_offload_pending(node_id) + self.assertEqual(core.get_write_through_pending_id(node_id), node_id) + self.assertFalse(core.is_external_cache_stored(node_id)) + + core.finish_external_linker_offload([node_id], node_id, success=True) + self.assertIsNone(core.get_write_through_pending_id(node_id)) + self.assertTrue(core.is_external_cache_stored(node_id)) + + def test_external_state_updates_are_atomic_and_path_scoped(self): + cfg = shared_cache_suite.CacheConfig( + page_size=1, kv_size=64, max_context_len=64 + ) + self.cfg = cfg + cache, allocator, req_to_token_pool = shared_cache_suite.build_fixture(cfg) + core = cache.tree_core + core.enable_external_cache_linker = True + + anchor = self._insert(cache, allocator, req_to_token_pool, [1]).last_device_node + middle = self._insert( + cache, allocator, req_to_token_pool, [1, 2] + ).last_device_node + tail = self._insert( + cache, allocator, req_to_token_pool, [1, 2, 3, 4] + ).last_device_node + unrelated = self._insert( + cache, allocator, req_to_token_pool, [9] + ).last_device_node + self.assertEqual(core.get_parent_node_id(middle), anchor) + self.assertEqual(core.get_parent_node_id(tail), middle) + + with self.assertRaisesRegex(RuntimeError, "not an ancestor"): + core.mark_external_cache_stored_path(tail, unrelated) + self.assertFalse(core.is_external_cache_stored(middle)) + self.assertFalse(core.is_external_cache_stored(tail)) + + core.mark_external_cache_stored_path(tail, anchor) + self.assertTrue(core.is_external_cache_stored(tail)) + self.assertTrue(core.is_external_cache_stored(middle)) + self.assertFalse(core.is_external_cache_stored(anchor)) + self.assertFalse(core.is_external_cache_stored(unrelated)) + with self.assertRaisesRegex(AssertionError, "invalid external offload state"): + core.mark_external_linker_offload_pending(tail) + + split_tail = self._insert( + cache, allocator, req_to_token_pool, [9, 10, 11] + ).last_device_node + self.assertEqual(core.get_parent_node_id(split_tail), unrelated) + core.mark_external_linker_offload_pending(split_tail) + self._insert(cache, allocator, req_to_token_pool, [9, 10]) + split_parent = core.get_parent_node_id(split_tail) + self.assertIsNotNone(split_parent) + self.assertNotEqual(split_parent, unrelated) + self.assertEqual(core.get_parent_node_id(split_parent), unrelated) + for node_id in (split_parent, split_tail): + self.assertEqual(core.get_write_through_pending_id(node_id), split_tail) + self.assertFalse(core.is_external_cache_stored(node_id)) + + independent = self._insert( + cache, allocator, req_to_token_pool, [20] + ).last_device_node + core.mark_external_linker_offload_pending(independent) + with self.assertRaisesRegex(AssertionError, "invalid external offload state"): + core.finish_external_linker_offload( + [independent, split_parent], independent, success=False + ) + self.assertEqual(core.get_write_through_pending_id(independent), independent) + for node_id in (split_parent, split_tail): + self.assertEqual(core.get_write_through_pending_id(node_id), split_tail) + self.assertFalse(core.is_external_cache_stored(node_id)) + + core.finish_external_linker_offload([independent], independent, success=False) + core.finish_external_linker_offload( + [split_parent, split_tail], split_tail, success=False + ) + for node_id in (split_parent, split_tail): + self.assertIsNone(core.get_write_through_pending_id(node_id)) + self.assertFalse(core.is_external_cache_stored(node_id)) + self.assertTrue(core.is_external_cache_stored(middle)) + self.assertTrue(core.is_external_cache_stored(tail)) + + def test_failed_offload_preserves_independently_confirmed_state(self): + cfg = shared_cache_suite.CacheConfig( + page_size=1, kv_size=64, max_context_len=64 + ) + self.cfg = cfg + cache, allocator, req_to_token_pool = shared_cache_suite.build_fixture(cfg) + core = cache.tree_core + core.enable_external_cache_linker = True + + anchor = self._insert(cache, allocator, req_to_token_pool, [1]).last_device_node + node_id = self._insert( + cache, allocator, req_to_token_pool, [1, 2] + ).last_device_node + core.mark_external_linker_offload_pending(node_id) + + core.mark_external_cache_stored_path(node_id, anchor) + self.assertEqual(core.get_write_through_pending_id(node_id), node_id) + self.assertTrue(core.is_external_cache_stored(node_id)) + + core.finish_external_linker_offload([node_id], node_id, success=False) + self.assertIsNone(core.get_write_through_pending_id(node_id)) + self.assertTrue(core.is_external_cache_stored(node_id)) + + +class TestUnifiedCacheLinkerRustBackend(TestUnifiedCacheLinkerPythonBackend): + tree_core_backend = "rust" + + +class TestUnifiedCacheLinkerTreeCoreRustBackend( + TestUnifiedCacheLinkerTreeCorePythonBackend +): + tree_core_backend = "rust" + + def test_restorable_prefix_intersects_sparse_rank_results(): remote_mask = torch.tensor([0, 0, 1, 0, 0], dtype=torch.int) @@ -127,11 +762,6 @@ def test_restorable_prefix_intersects_sparse_rank_results(): def test_async_offload_pins_node_until_completion(): - class _Component: - def build_external_linker_transfer(self, phase, node, keys): - assert phase == LinkerTransferPhase.OFFLOAD - return PoolTransfer(name=PoolName.KV, keys=["page"]) - linker = _FakeLinker() lock_params = object() locks = [] @@ -148,23 +778,17 @@ def test_async_offload_pins_node_until_completion(): write_through_pending_id=None, ) cache = _cache_for_wrapper( - tree_core=SimpleNamespace( - enable_external_cache_linker=False, - mark_write_through_pending=lambda node_ids, ack_id: ( - setattr(node, "write_through_pending_id", ack_id) or list(node_ids) - ), - ), - _components_tuple=(_Component(),), + tree_core=_FakeExternalTreeCore({node_id: node}), inc_lock_ref=inc_lock_ref, dec_lock_ref=lambda node, params: unlocks.append((node, params)), - resolve_node_handle=lambda value: node if value == node_id else None, ) wrapper = UnifiedCacheLinkerWrapper(cache, linker) wrapper.offload_nodes([node_id]) assert locks == [node_id] - assert node.external_cache_stored + assert not node.external_cache_stored + assert node.write_through_pending_id == node_id assert not unlocks linker.completed_offloads.append(False) @@ -172,9 +796,28 @@ def test_async_offload_pins_node_until_completion(): wrapper.commit_completed_offloads(completed) assert not node.external_cache_stored + assert node.write_through_pending_id is None assert unlocks == [(node_id, lock_params)] +def test_offload_skips_node_already_stored_by_tree_core(): + linker = _FakeLinker() + node = SimpleNamespace( + id=7, + external_cache_stored=True, + write_through_pending_id=None, + ) + cache = _cache_for_wrapper( + tree_core=_FakeExternalTreeCore({node.id: node}), + inc_lock_ref=lambda node_id: pytest.fail("stored node must not be locked"), + ) + wrapper = UnifiedCacheLinkerWrapper(cache, linker) + + wrapper.offload_nodes([node.id]) + + assert linker.queued_offloads == [] + + def test_async_load_pins_node_until_completion(): linker = _FakeLinker() lock_params = object() @@ -224,10 +867,6 @@ def test_release_request_cancels_queued_load(): def test_failed_offload_rolls_back_split_fragments(): - class _Component: - def build_external_linker_transfer(self, phase, node, keys): - return PoolTransfer(name=PoolName.KV, keys=["page"]) - linker = _FakeLinker() lock_params = object() unlocks = [] @@ -243,20 +882,10 @@ def test_failed_offload_rolls_back_split_fragments(): ) nodes = {child.id: child, parent.id: parent} - def mark_pending(node_ids, ack_id): - for node_id in node_ids: - nodes[node_id].write_through_pending_id = ack_id - return list(node_ids) - cache = _cache_for_wrapper( - tree_core=SimpleNamespace( - enable_external_cache_linker=False, - mark_write_through_pending=mark_pending, - ), - _components_tuple=(_Component(),), + tree_core=_FakeExternalTreeCore(nodes), inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=lambda: lock_params), dec_lock_ref=lambda node_id, params: unlocks.append((node_id, params)), - resolve_node_handle=nodes.__getitem__, ) wrapper = UnifiedCacheLinkerWrapper(cache, linker) wrapper.offload_nodes([child.id]) @@ -299,10 +928,6 @@ def test_split_action_retargets_pending_external_offload(): def test_reset_quiesces_backend_before_releasing_pending_locks(): - class _Component: - def build_external_linker_transfer(self, phase, node, keys): - return PoolTransfer(name=PoolName.KV, keys=["page"]) - events = [] class _QuiescentFakeLinker(_FakeLinker): @@ -317,16 +942,9 @@ def test_reset_quiesces_backend_before_releasing_pending_locks(): write_through_pending_id=None, ) cache = _cache_for_wrapper( - tree_core=SimpleNamespace( - enable_external_cache_linker=False, - mark_write_through_pending=lambda node_ids, ack_id: ( - setattr(node, "write_through_pending_id", ack_id) or list(node_ids) - ), - ), - _components_tuple=(_Component(),), + tree_core=_FakeExternalTreeCore({node.id: node}), inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=object), dec_lock_ref=lambda node_id, params: events.append(("unlock", node_id)), - resolve_node_handle=lambda node_id: node, ) wrapper = UnifiedCacheLinkerWrapper(cache, linker) wrapper._queue_load("rid", node.id, [object()]) diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 5a04e4da7..2fed5a542 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -8117,6 +8117,7 @@ class _InsertWalkSuite(CustomTestCase): _rid = 0 _make_req = UnifiedRadixCacheSuite._make_req + _apply_match_to_req = UnifiedRadixCacheSuite._apply_match_to_req _alloc = UnifiedRadixCacheSuite._alloc _insert = UnifiedRadixCacheSuite._insert _init_hicache = UnifiedRadixCacheSuite._init_hicache diff --git a/test/registered/unit/mem_cache/unified_tree_core_inspection_interface.py b/test/registered/unit/mem_cache/unified_tree_core_inspection_interface.py index 1567b9a15..4749faf04 100644 --- a/test/registered/unit/mem_cache/unified_tree_core_inspection_interface.py +++ b/test/registered/unit/mem_cache/unified_tree_core_inspection_interface.py @@ -88,6 +88,11 @@ class UnifiedTreeCoreInspectionInterface(UnifiedTreeCoreInterface): """The node's pending write-through id, if any.""" ... + @abstractmethod + def is_external_cache_stored(self, node_id: NodeId) -> bool: + """Whether the node is known to be stored in the external cache.""" + ... + @abstractmethod def is_node_in_device_lru( self, node_id: NodeId, component_type: ComponentType diff --git a/test/registered/unit/mem_cache/unified_tree_core_inspector.py b/test/registered/unit/mem_cache/unified_tree_core_inspector.py index 529b1ee50..c13d145ed 100644 --- a/test/registered/unit/mem_cache/unified_tree_core_inspector.py +++ b/test/registered/unit/mem_cache/unified_tree_core_inspector.py @@ -81,6 +81,10 @@ class UnifiedTreeCoreInspector(UnifiedTreeCore, UnifiedTreeCoreInspectionInterfa """The node's pending write-through id, if any.""" return self.node_by_id(node_id).write_through_pending_id + def is_external_cache_stored(self, node_id: NodeId) -> bool: + """Whether the node is known to be stored in the external cache.""" + return self.node_by_id(node_id).external_cache_stored + def is_node_in_device_lru( self, node_id: NodeId, component_type: ComponentType ) -> bool: