[Rust TreeCore] Harden runtime and CI parity (#37303)

This commit is contained in:
Jialin Ouyang
2026-09-09 10:40:22 +08:00
committed by GitHub
parent e54ff1efb9
commit 7a464a7014
22 changed files with 2390 additions and 1137 deletions
@@ -7,6 +7,7 @@ from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
is_in_amd_ci,
popen_launch_server,
terminate_and_kill_process_tree,
unified_radix_tree_server_env,
@@ -49,6 +50,7 @@ class TestUnifiedFullRadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
terminate_and_kill_process_tree(cls.process, wait_timeout=60)
@unittest.skipIf(is_in_amd_ci(), "Rust TreeCore is not packaged in AMD CI")
class TestRustUnifiedFullRadixCache(TestUnifiedFullRadixCache):
tree_core_backend = "rust"
+38 -1
View File
@@ -162,6 +162,7 @@ crate-type = ["cdylib"]
self.assertEqual(crate.library, "demo_extension")
self.assertEqual(crate.python_module, "demo._core")
self.assertEqual(crate.features, ("python",))
self.assertEqual(crate.source_inputs, ())
with self.assertRaisesRegex(
ModuleNotFoundError, r"declared modules: \['demo\._core'\]"
@@ -208,6 +209,36 @@ crate-type = ["cdylib"]
inspection.target_fingerprint,
)
def test_fingerprint_covers_declared_external_source_inputs(self):
with TemporaryDirectory() as directory:
root = Path(directory)
workspace = self._workspace(root)
proto = root / "proto/demo.proto"
proto.parent.mkdir()
proto.write_text("message Demo {}\n", encoding="utf-8")
manifest = workspace / "demo/Cargo.toml"
manifest.write_text(
manifest.read_text(encoding="utf-8").replace(
'features = ["python"]',
'features = ["python"]\nsource-inputs = ["../../proto"]',
),
encoding="utf-8",
)
crate = rust_extension._discover_crate(workspace, "demo._core")
self.assertEqual(crate.source_inputs, (proto.parent.resolve(),))
with mock.patch.object(
rust_extension,
"_command_version",
side_effect=lambda command, *args, **kwargs: f"{command} 1.0",
):
first = rust_extension._build_context(crate)
proto.write_text("message Changed {}\n", encoding="utf-8")
changed = rust_extension._build_context(crate)
self.assertNotEqual(first.fingerprint, changed.fingerprint)
self.assertEqual(first.target_fingerprint, changed.target_fingerprint)
def test_auto_builds_once_then_uses_cache(self):
with TemporaryDirectory() as directory:
root = Path(directory)
@@ -527,30 +558,35 @@ crate-type = ["cdylib"]
self.assertNotIn(module_name, sys.modules)
def test_checked_in_crates_are_discovered_from_wheel_metadata(self):
for python_module, package, library, features in (
grpc_proto = (rust_extension._RUST_WORKSPACE.parent / "proto").resolve()
for python_module, package, library, features, source_inputs in (
(
"sglang.srt.rust_extensions._server",
"sglang-server",
"sglang_server",
(),
(),
),
(
"sglang.srt.rust_extensions._grpc",
"sglang-grpc",
"sglang_grpc_core",
(),
(grpc_proto,),
),
(
"sglang.srt.rust_extensions._multimodal",
"sglang-mm",
"sglang_mm_core",
("python", "parallel"),
(),
),
(
"sglang.srt.mem_cache.rust_tree_core.mem_cache",
"sglang-radix-tree",
"mem_cache",
("python-extension",),
(),
),
):
crate = rust_extension._discover_crate(
@@ -559,6 +595,7 @@ crate-type = ["cdylib"]
self.assertEqual(crate.package, package)
self.assertEqual(crate.library, library)
self.assertEqual(crate.features, features)
self.assertEqual(crate.source_inputs, source_inputs)
if __name__ == "__main__":
@@ -29,6 +29,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
InsertParams,
InsertResult,
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.hicache_storage import (
@@ -173,6 +174,51 @@ def test_stale_handle_reads_raise_key_error_without_poisoning_the_core():
assert core.is_root(live_root)
def test_stale_match_finalizer_handles_raise_key_error_without_poisoning_the_core():
from rust_unified_tree_core_inspector import RustUnifiedTreeCoreInspector
core = RustUnifiedTreeCoreInspector(
CacheInitParams(
disable=False,
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
page_size=1,
tree_components=(ComponentType.FULL,),
)
)
stale_root = core.root_node_handle()
core.reset()
live_root = core.root_node_handle()
result = MatchResult(
device_indices=torch.empty(0, dtype=torch.int64),
last_device_node=live_root,
last_host_node=live_root,
best_match_node=live_root,
)
params = MatchPrefixParams(key=_key([]))
for field in ("last_device_node", "last_host_node", "best_match_node"):
with pytest.raises(KeyError) as exc_info:
core.finalize_component_match_result(
ComponentType.FULL,
result._replace(**{field: stale_root}),
params,
value_chunks=[],
best_value_len=0,
)
assert exc_info.value.args == (stale_root,), field
assert core.is_root(live_root), field
finalized = core.finalize_component_match_result(
ComponentType.FULL,
result,
params,
value_chunks=[],
best_value_len=0,
)
assert finalized.best_match_node == live_root
def test_stale_handle_operations_raise_key_error_without_poisoning_the_core():
from sglang.srt.mem_cache.unified_cache.components import CacheTransferPhase
@@ -180,15 +226,122 @@ def test_stale_handle_operations_raise_key_error_without_poisoning_the_core():
stale_root = core.root_node_handle()
core.reset()
live_root = core.root_node_handle()
empty = torch.empty(0, dtype=torch.int64)
operations = (
lambda: core.demote(stale_root),
lambda: core.build_hicache_transfers(
operations = {
"inc_lock_ref": lambda: core.inc_lock_ref(stale_root),
"dec_lock_ref": lambda: core.dec_lock_ref(stale_root),
"dec_swa_lock_only": lambda: core.dec_swa_lock_only(stale_root, None),
"evict_device_leaf": lambda: core.evict_device_leaf(stale_root, False),
"drop_subtree_no_host": lambda: core.drop_subtree_no_host(stale_root),
"demote": lambda: core.demote(stale_root),
"is_full_device_evicted": lambda: core.is_full_device_evicted(stale_root),
"collect_full_device_indices/from": lambda: core.collect_full_device_indices(
stale_root, live_root
),
"collect_full_device_indices/until": lambda: core.collect_full_device_indices(
live_root, stale_root
),
"insert_host": lambda: core.insert_host(
stale_root, _key([1]), empty, ["0" * 64]
),
"build_backup_spec": lambda: core.build_backup_spec(stale_root),
"build_storage_backup_spec": lambda: core.build_storage_backup_spec(
stale_root, False
),
"build_hicache_transfers": lambda: core.build_hicache_transfers(
ComponentType.FULL, stale_root, CacheTransferPhase.BACKUP_STORAGE
),
lambda: core.build_load_back_spec(stale_root),
lambda: core.get_hash_values(stale_root),
lambda: core.dfs_weight_order([stale_root]),
"commit_backup": lambda: core.commit_backup(stale_root, empty, {}),
"commit_hicache_transfers": lambda: core.commit_hicache_transfers(
stale_root,
CacheTransferPhase.BACKUP_HOST,
{},
cache_actions=[],
),
"commit_load_back": lambda: core.commit_load_back(
stale_root, empty, PoolTransfer(name=PoolName.KV), {}
),
"build_load_back_spec": lambda: core.build_load_back_spec(stale_root),
"evict_excess_path_states": lambda: core.evict_excess_path_states(
stale_root, {}, {}
),
"inc_host_lock_ref": lambda: core.inc_host_lock_ref(stale_root),
"dec_host_lock_ref": lambda: core.dec_host_lock_ref(stale_root),
"mark_write_through_pending": lambda: core.mark_write_through_pending(
[stale_root], stale_root
),
"finish_write_through": lambda: core.finish_write_through(
[stale_root], stale_root
),
"finish_load_back": lambda: core.finish_load_back(stale_root),
"get_component_device_value": lambda: core.get_component_device_value(
stale_root, ComponentType.FULL
),
"component_has_host_value_only": lambda: core.component_has_host_value_only(
stale_root, ComponentType.FULL
),
"get_hash_values": lambda: core.get_hash_values(stale_root),
"dfs_weight_order": lambda: core.dfs_weight_order([stale_root]),
}
for name, operation in operations.items():
with pytest.raises(KeyError) as exc_info:
operation()
assert exc_info.value.args == (stale_root,), name
assert core.is_root(live_root), name
def test_stale_handles_nested_in_transfer_results_do_not_poison_the_core():
from sglang.srt.mem_cache.unified_cache.components import CacheTransferPhase
core = _tree_core()
stale_root = core.root_node_handle()
core.reset()
live_root = core.root_node_handle()
stale_transfer = PoolTransfer(name=PoolName.KV, nodes_to_load=[stale_root])
operations = (
lambda: core.commit_hicache_transfers(
live_root,
CacheTransferPhase.LOAD_BACK,
{ComponentType.FULL: [stale_transfer]},
cache_actions=[],
),
lambda: core.commit_hicache_transfers(
live_root,
CacheTransferPhase.PREFETCH,
{},
cache_actions=[],
insert_result=InsertResult(prefix_len=0, inserted_host_node=stale_root),
),
lambda: core.commit_load_back(
live_root,
torch.empty(0, dtype=torch.int64),
stale_transfer,
{},
),
)
for operation in operations:
with pytest.raises(KeyError) as exc_info:
operation()
assert exc_info.value.args == (stale_root,)
assert core.is_root(live_root)
def test_stale_handle_component_access_does_not_poison_the_core():
core = _tree_core(
tree_components=(ComponentType.FULL, ComponentType.SWA),
sliding_window_size=8,
)
stale_root = core.root_node_handle()
core.reset()
live_root = core.root_node_handle()
operations = (
lambda: core.set_component_device_value(
stale_root, ComponentType.SWA, torch.empty(0, dtype=torch.int64)
),
lambda: core.get_component_device_value(stale_root, ComponentType.SWA),
)
for operation in operations:
with pytest.raises(KeyError) as exc_info:
@@ -1261,6 +1414,10 @@ def test_buffer_backup_snapshot_round_trips_and_detects_a_split():
)
assert core.validate_buffer_backup(leaf, len(snapshot.key)) is None
core.reset()
assert core.snapshot_buffer_backup(leaf, pass_prefix_keys=True) is None
assert core.validate_buffer_backup(leaf, len(snapshot.key)) is None
def test_buffer_backup_snapshot_preserves_bigram_keys():
core = _tree_core(is_eagle=True)
@@ -1996,5 +2153,122 @@ def test_bigram_insert_value_shorter_than_the_bigram_count_raises():
)
def test_stale_inspection_handles_raise_key_error_or_report_absence():
from rust_unified_tree_core_inspector import RustUnifiedTreeCoreInspector
from sglang.srt.mem_cache.unified_cache.components import EvictLayer
core = RustUnifiedTreeCoreInspector(
CacheInitParams(
disable=False,
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
page_size=1,
tree_components=(ComponentType.FULL,),
)
)
stale_root = core.root_node_handle()
core.reset()
live_root = core.root_node_handle()
operations = {
"get_parent_node_id": lambda: core.get_parent_node_id(stale_root),
"get_child_node_ids": lambda: core.get_child_node_ids(stale_root),
"get_node_key_length": lambda: core.get_node_key_length(stale_root),
"get_node_token_ids": lambda: core.get_node_token_ids(stale_root),
"is_node_key_bigram": lambda: core.is_node_key_bigram(stale_root),
"get_component_host_value": lambda: core.get_component_host_value(
stale_root, ComponentType.FULL
),
"get_component_device_lock_ref": lambda: core.get_component_device_lock_ref(
stale_root, ComponentType.FULL
),
"get_node_hit_count": lambda: core.get_node_hit_count(stale_root),
"get_write_through_pending_id": lambda: core.get_write_through_pending_id(
stale_root
),
"is_node_in_device_lru": lambda: core.is_node_in_device_lru(
stale_root, ComponentType.FULL
),
"is_node_in_host_lru": lambda: core.is_node_in_host_lru(
stale_root, ComponentType.FULL
),
"is_device_leaf": lambda: core.is_device_leaf(stale_root),
"set_node_hash_values": lambda: core.set_node_hash_values(stale_root, None),
"set_component_device_value_raw": lambda: core.set_component_device_value_raw(
stale_root, ComponentType.FULL, None
),
"set_component_host_value_raw": lambda: core.set_component_host_value_raw(
stale_root, ComponentType.FULL, None
),
"set_component_device_lock_ref": lambda: core.set_component_device_lock_ref(
stale_root, ComponentType.FULL, 0
),
"remove_node_from_device_lru": lambda: core.remove_node_from_device_lru(
stale_root, ComponentType.FULL
),
"insert_node_into_host_lru": lambda: core.insert_node_into_host_lru(
stale_root, ComponentType.FULL
),
"update_duplicate_tracking": lambda: core.update_duplicate_tracking(stale_root),
"evict_component": lambda: core.evict_component(
stale_root, ComponentType.FULL, EvictLayer.DEVICE
),
"validate_cascade_evict": lambda: core.validate_cascade_evict(
stale_root, ComponentType.FULL, EvictLayer.DEVICE
),
"cleanup_tombstone_ancestors": lambda: core.cleanup_tombstone_ancestors(
stale_root
),
"build_backup_node_ids": lambda: core.build_backup_node_ids(stale_root),
}
for name, operation in operations.items():
with pytest.raises(KeyError) as exc_info:
operation()
assert exc_info.value.args == (stale_root,), name
assert core.is_root(live_root), name
disabled_component_operations = {
"get_component_host_value": lambda: core.get_component_host_value(
stale_root, ComponentType.SWA
),
"get_component_device_lock_ref": lambda: core.get_component_device_lock_ref(
stale_root, ComponentType.SWA
),
"set_component_device_value_raw": lambda: core.set_component_device_value_raw(
stale_root, ComponentType.SWA, None
),
"set_component_host_value_raw": lambda: core.set_component_host_value_raw(
stale_root, ComponentType.SWA, None
),
"set_component_device_lock_ref": lambda: core.set_component_device_lock_ref(
stale_root, ComponentType.SWA, 0
),
"remove_node_from_device_lru": lambda: core.remove_node_from_device_lru(
stale_root, ComponentType.SWA
),
"insert_node_into_host_lru": lambda: core.insert_node_into_host_lru(
stale_root, ComponentType.SWA
),
"evict_component": lambda: core.evict_component(
stale_root, ComponentType.SWA, EvictLayer.DEVICE
),
"validate_cascade_evict": lambda: core.validate_cascade_evict(
stale_root, ComponentType.SWA, EvictLayer.DEVICE
),
}
for name, operation in disabled_component_operations.items():
with pytest.raises(KeyError) as exc_info:
operation()
assert exc_info.value.args == (stale_root,), name
assert core.is_root(live_root), name
assert not core.contains_node(stale_root)
assert not core.is_device_evictable_leaf(stale_root)
assert not core.is_host_evictable_leaf(stale_root)
assert not core.is_node_in_device_lru(stale_root, ComponentType.SWA)
assert not core.is_node_in_host_lru(stale_root, ComponentType.SWA)
if __name__ == "__main__":
sys.exit(pytest.main([__file__]))
@@ -8065,6 +8065,61 @@ class _InsertWalkSuite(CustomTestCase):
)
@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA")
class TestUnifiedTreeCoreSWAPrefetchBackends(_InsertWalkSuite):
cfg = CacheConfig(
components=(ComponentType.FULL, ComponentType.SWA), sliding_window_size=4
)
def test_mid_tree_shortened_swa_prefetch_is_released(self):
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
prefix = [1, 2]
self._insert(cache, allocator, req_to_token_pool, prefix)
anchor = cache.match_prefix(
MatchPrefixParams(key=RadixKey(array("q", prefix)))
).last_device_node
cache.tree_core.is_write_back = True
suffix = [3, 4]
insert_result = cache.tree_core.insert_host(
anchor,
RadixKey(array("q", suffix)),
torch.tensor([100, 101], dtype=torch.int64),
["h3", "h4"],
)
self.assertIsNotNone(insert_result.inserted_host_node)
swa_host_indices = torch.tensor([30, 31], dtype=torch.int64)
actions = []
cache.tree_core.commit_hicache_transfers(
anchor,
CacheTransferPhase.PREFETCH,
{
ComponentType.SWA: [
PoolTransfer(
name=PoolName.SWA,
host_indices=swa_host_indices,
)
]
},
cache_actions=actions,
insert_result=insert_result,
pool_storage_result=PoolTransferResult(
kv_hit_pages=len(suffix),
extra_pool_hit_pages={PoolName.SWA: len(suffix)},
),
)
self.assertIsNone(
_host_value(cache, insert_result.inserted_host_node, ComponentType.SWA)
)
self.assertEqual(len(actions), 1)
self.assertIsInstance(actions[0], FreeComponentHostSlot)
self.assertEqual(actions[0].component_type, ComponentType.SWA)
self.assertEqual(len(actions[0].host_indices), 1)
self.assertTrue(torch.equal(actions[0].host_indices[0], swa_host_indices))
@unittest.skipUnless(torch.cuda.is_available(), "cache fixtures need CUDA")
class TestResumableInsertWalk(_InsertWalkSuite):
cfg = CacheConfig()