""" Unit tests for RadixKey zero-copy view semantics. RadixKey slicing returns O(1) views over a shared backing list instead of copying. These tests pin down the semantics that the radix tree relies on: - Slicing/iteration/indexing behave exactly like list slicing. - Views share the backing list; compacted() detaches them. - Key matching and child-key derivation are view-position independent. - Keys stored in tree nodes are always compact (never views), so lock-ref walks, eviction, splits, and controller-thread reads never observe a key that pins a transient backing list. Usage: python -m pytest test_radix_key_views.py -v """ import sys import unittest.mock for _mod in ("sgl_kernel", "sgl_kernel.kvcacheio"): if _mod not in sys.modules: sys.modules[_mod] = unittest.mock.MagicMock() from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci # CPU-based unit test, runs quickly on any GPU runner register_cuda_ci(est_time=5, suite="stage-b-test-1-gpu-small") register_amd_ci(est_time=5, suite="stage-b-test-1-gpu-small-amd") import random import unittest import torch from sglang.srt.mem_cache.base_prefix_cache import InsertParams, MatchPrefixParams from sglang.srt.mem_cache.radix_cache import ( RadixCache, RadixKey, _key_match_page_size1, _key_match_paged, get_child_key, maybe_bigram_convert, ) class TestBigramConversionEquivalence(unittest.TestCase): def test_zip_impl_matches_index_comprehension(self): from sglang.srt.mem_cache.utils import _python_convert_to_bigram_key rng = random.Random(7) for n in (0, 1, 2, 3, 17, 256): tokens = [rng.randrange(1000) for _ in range(n)] want = [(tokens[i], tokens[i + 1]) for i in range(max(len(tokens) - 1, 0))] self.assertEqual(_python_convert_to_bigram_key(tokens), want) # Already-converted (tuple) input is returned unchanged. pairs = [(1, 2), (2, 3)] self.assertEqual(_python_convert_to_bigram_key(pairs), pairs) class TestRadixKeyViewSemantics(unittest.TestCase): def test_slice_is_view_sharing_backing(self): backing = list(range(100)) key = RadixKey(backing) view = key[10:50] self.assertIs(view._tokens, key._tokens) self.assertEqual(view.token_ids, backing[10:50]) self.assertEqual(len(view), 40) def test_nested_slicing_matches_list_semantics(self): rng = random.Random(0) tokens = [rng.randrange(1000) for _ in range(257)] key = RadixKey(tokens) ref = list(tokens) for _ in range(200): if len(ref) == 0: break a = rng.randrange(0, len(ref) + 1) b = rng.randrange(a, len(ref) + 1) key = key[a:b] ref = ref[a:b] self.assertEqual(key.token_ids, ref) self.assertEqual(len(key), len(ref)) self.assertEqual(list(key), ref) def test_open_ended_and_negative_slices(self): tokens = list(range(20)) key = RadixKey(tokens) self.assertEqual(key[5:].token_ids, tokens[5:]) self.assertEqual(key[:7].token_ids, tokens[:7]) self.assertEqual(key[-4:].token_ids, tokens[-4:]) self.assertEqual(key[3:1].token_ids, []) def test_int_indexing(self): tokens = [7, 8, 9, 10] key = RadixKey(tokens)[1:] item = key[0] self.assertEqual(item.token_ids, [8]) self.assertEqual(key[-1].token_ids, [10]) def test_strided_slicing_rejected(self): key = RadixKey(list(range(10))) with self.assertRaises(ValueError): key[::2] def test_compacted_detaches_views_and_is_identity_for_full_keys(self): backing = list(range(50)) key = RadixKey(backing, extra_key="lora1", is_bigram=False) self.assertIs(key.compacted(), key) view = key[10:30] compact = view.compacted() self.assertIsNot(compact._tokens, backing) self.assertEqual(compact.token_ids, backing[10:30]) self.assertEqual(compact.extra_key, "lora1") # Mutating the original backing list must not affect the compacted key. backing[15] = -1 self.assertNotIn(-1, compact.token_ids) def test_token_ids_setter_resets_view(self): key = RadixKey(list(range(30)))[5:25] key.token_ids = [1, 2, 3] self.assertEqual(len(key), 3) self.assertEqual(key.token_ids, [1, 2, 3]) def test_length_frozen_at_construction(self): backing = [1, 2, 3] key = RadixKey(backing) backing.append(4) self.assertEqual(len(key), 3) self.assertEqual(key.token_ids, [1, 2, 3]) def test_extra_key_and_bigram_preserved_through_slices(self): key = RadixKey([(1, 2), (2, 3), (3, 4)], extra_key="salt", is_bigram=True) view = key[1:] self.assertEqual(view.extra_key, "salt") self.assertTrue(view.is_bigram) self.assertTrue(view.compacted().is_bigram) def test_bigram_convert_on_view(self): key = RadixKey(list(range(10)))[2:8] converted, _ = maybe_bigram_convert(True, key) self.assertTrue(converted.is_bigram) self.assertEqual( converted.token_ids[0][0] if converted.token_ids else None, 2, ) def test_key_match_is_view_position_independent(self): rng = random.Random(1) base = [rng.randrange(50) for _ in range(512)] other = list(base[:300]) + [999] + base[301:] for page_size in (1, 16, 64): for off in (0, 1, 63, 128): k0 = RadixKey([0] * off + base)[off:] k1 = RadixKey(other) k0_ref = RadixKey(list(base)) if page_size == 1: got = _key_match_page_size1(k0, k1) want = _key_match_page_size1(k0_ref, k1) else: got = _key_match_paged(k0, k1, page_size) want = _key_match_paged(k0_ref, k1, page_size) self.assertEqual(got, want, f"{page_size=} {off=}") def test_get_child_key_on_views(self): tokens = list(range(200)) key = RadixKey(tokens, extra_key="e") view = key[64:] self.assertEqual(get_child_key(view, 64), get_child_key(view.compacted(), 64)) self.assertEqual(get_child_key(view, 64), ("e", tuple(tokens[64:128]))) self.assertEqual(get_child_key(view, 1), ("e", 64)) def test_tree_node_keys_are_always_compact(self): """Node-stored keys must never be views: a view would pin the transient request key's backing list, and lock/update paths (inc_lock_ref, eviction, splits) would read keys whose backing outlives the match.""" rng = random.Random(2) cache = RadixCache.create_simulated(page_size=4) seqs = [] base = [rng.randrange(8) for _ in range(256)] for _ in range(40): # Shared prefixes of random length force splits at random points. cut = rng.randrange(8, 256, 4) seq = base[:cut] + [rng.randrange(8) for _ in range(rng.randrange(4, 64))] seqs.append(seq) cache.insert( InsertParams( key=RadixKey(seq), value=torch.arange(len(seq), dtype=torch.int64), ) ) cache.match_prefix(MatchPrefixParams(key=RadixKey(list(seq)))) stack = [cache.root_node] checked = 0 while stack: node = stack.pop() stack.extend(node.children.values()) if node is cache.root_node or node.key is None: continue checked += 1 self.assertEqual(node.key._start, 0, f"node {node.id} key is a view") self.assertEqual( node.key._stop, len(node.key._tokens), f"node {node.id} key is a view", ) self.assertGreater(checked, 10) def test_match_and_insert_equivalence_random(self): """Differential test: tree behavior identical for fresh keys vs views.""" rng = random.Random(3) cache = RadixCache.create_simulated(page_size=2) for _ in range(50): seq = [rng.randrange(4) for _ in range(rng.randrange(2, 40, 2))] cache.insert( InsertParams( key=RadixKey(seq), value=torch.arange(len(seq), dtype=torch.int64), ) ) probe = [rng.randrange(4) for _ in range(rng.randrange(2, 40, 2))] res_fresh = cache.match_prefix(MatchPrefixParams(key=RadixKey(list(probe)))) padded = RadixKey([9] * 6 + list(probe))[6:] res_view = cache.match_prefix(MatchPrefixParams(key=padded)) self.assertTrue( torch.equal(res_fresh.device_indices, res_view.device_indices) ) if __name__ == "__main__": unittest.main()