feat(agent sessions): attribute stored KV cache blocks to sessions (#37482)

Signed-off-by: Ishan Dhanani <ishandhanani@gmail.com>
This commit is contained in:
ishandhanani
2026-09-14 21:34:49 -07:00
committed by GitHub
parent c9fbe5f655
commit 3f871a246c
22 changed files with 905 additions and 451 deletions
@@ -612,6 +612,7 @@ class _GraftReq:
self.swa_prefix_lock_released = False
self.finished_reason = None
self.session = None
self.session_id = None
def get_fill_ids(self):
return array("q", self.fill_ids)
@@ -34,8 +34,6 @@ from sglang.srt.disaggregation.kv_events import (
AllBlocksCleared,
BlockRemoved,
BlockStored,
BlockStoredMetadata,
BlockStoredWithMetadata,
StorageMedium,
)
from sglang.srt.managers.schedule_batch import ReqKvInfo
@@ -66,20 +64,17 @@ class TestKVCacheEventQueue(unittest.TestCase):
medium: StorageMedium = StorageMedium.GPU,
lora_id: int | None = None,
cache_salt: str | None = None,
session_id: str | None = None,
) -> BlockStored:
event_args = dict(
return BlockStored(
block_hashes=[block_hash],
parent_block_hash=parent_block_hash,
token_ids=[block_hash, block_hash + 1][:block_size],
block_size=block_size,
lora_id=lora_id,
medium=medium,
)
if cache_salt is None:
return BlockStored(**event_args)
return BlockStoredWithMetadata(
**event_args,
metadata=BlockStoredMetadata(cache_salt=cache_salt),
cache_salt=cache_salt,
session_id=session_id,
)
def test_enqueue_coalesces_compatible_stores(self):
@@ -133,6 +128,11 @@ class TestKVCacheEventQueue(unittest.TestCase):
queue.enqueue(self._store(2, 1, cache_salt="tenant-b"))
self.assertEqual(len(queue.take()), 2)
queue = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
queue.enqueue(self._store(1, None, session_id="session-a"))
queue.enqueue(self._store(2, 1, session_id="session-b"))
self.assertEqual(len(queue.take()), 2)
class TestRadixKey(unittest.TestCase):
"""Test cases for RadixKey class."""
@@ -781,7 +781,7 @@ class TestRadixCache(CustomTestCase):
removed = [event for event in events if isinstance(event, BlockRemoved)]
self.assertEqual(len(stored), 1)
self.assertEqual(stored[0].metadata.cache_salt, "tenant-a")
self.assertEqual(stored[0].cache_salt, "tenant-a")
self.assertEqual(stored[0].parent_block_hash, None)
self.assertEqual(len(stored[0].block_hashes), 2)
self.assertEqual(removed[0].block_hashes, stored[0].block_hashes)
@@ -20,8 +20,6 @@ from sglang.srt.disaggregation.kv_events import (
AllBlocksCleared,
BlockRemoved,
BlockStored,
BlockStoredMetadata,
BlockStoredWithMetadata,
StorageMedium,
)
from sglang.srt.environ import envs
@@ -971,14 +969,14 @@ def test_salted_events_match_python_hash_and_metadata_contract():
for value in mem_cache.get_hash_str(array("q", [1, 2, 7, 8]), seed, 2)
]
assert core.take_events() == [
BlockStoredWithMetadata(
BlockStored(
block_hashes=hashes,
parent_block_hash=None,
token_ids=[1, 2, 7, 8],
block_size=2,
lora_id=None,
medium=StorageMedium.GPU,
metadata=BlockStoredMetadata(cache_salt="tenant-a"),
cache_salt="tenant-a",
)
]
@@ -1012,14 +1010,14 @@ def test_salted_eagle_events_match_the_bigram_hash_contract():
for value in mem_cache.get_hash_str(raw_tokens, seed, 2, is_bigram=True)
]
assert core.take_events() == [
BlockStoredWithMetadata(
BlockStored(
block_hashes=hashes,
parent_block_hash=None,
token_ids=[(1, 2), (2, 3), (3, 4), (4, 5)],
block_size=2,
lora_id=None,
medium=StorageMedium.GPU,
metadata=BlockStoredMetadata(cache_salt="tenant-a"),
cache_salt="tenant-a",
)
]
@@ -22,7 +22,6 @@ from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.disaggregation.kv_events import (
BlockRemoved,
BlockStored,
BlockStoredWithMetadata,
StorageMedium,
)
from sglang.srt.environ import envs
@@ -1012,11 +1011,18 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
*,
extra_key=None,
cache_salt=None,
session_id=None,
):
key = RadixKey(array("q", tokens), extra_key=extra_key, cache_salt=cache_salt)
value = allocator.alloc(len(tokens))
self.assertIsNotNone(value)
return cache.insert(InsertParams(key=key, value=value[: len(key)]))
return cache.insert(
InsertParams(
key=key,
value=value[: len(key)],
session_id=session_id,
)
)
def _stored_events(self, cache, medium=None):
events = [e for e in cache.take_events() if isinstance(e, BlockStored)]
@@ -1096,8 +1102,7 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
self._insert(cache, allocator, seq, cache_salt="tenant-a")
stored = self._stored_events(cache, StorageMedium.GPU)
self.assertEqual(len(stored), 1)
self.assertIsInstance(stored[0], BlockStoredWithMetadata)
self.assertEqual(stored[0].metadata.cache_salt, "tenant-a")
self.assertEqual(stored[0].cache_salt, "tenant-a")
salted_hashes = self._event_hashes(stored)
cache.evict(EvictParams(num_tokens=len(seq)))
@@ -1115,6 +1120,73 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
salted_hashes,
)
def test_session_id_is_attributed_without_changing_block_hash(self):
cache_a, allocator_a, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
cache_a.take_events()
self._insert(cache_a, allocator_a, [1, 2, 3, 4], session_id="session-a")
stored_a = self._stored_events(cache_a, StorageMedium.GPU)
self.assertEqual(len(stored_a), 1)
self.assertEqual(stored_a[0].session_id, "session-a")
cache_b, allocator_b, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
cache_b.take_events()
self._insert(cache_b, allocator_b, [1, 2, 3, 4], session_id="session-b")
stored_b = self._stored_events(cache_b, StorageMedium.GPU)
self.assertEqual(stored_b[0].session_id, "session-b")
self.assertEqual(self._event_hashes(stored_a), self._event_hashes(stored_b))
def test_shared_prefix_hit_is_quiet_and_divergent_tails_are_attributed(self):
cache, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
cache.take_events()
shared_prefix = [1, 2, 3, 4]
self._insert(cache, allocator, shared_prefix, session_id="session-a")
initial = self._stored_events(cache, StorageMedium.GPU)
self.assertEqual(len(initial), 1)
shared_parent = initial[0].block_hashes[-1]
self._insert(cache, allocator, shared_prefix, session_id="session-b")
self.assertEqual(self._stored_events(cache, StorageMedium.GPU), [])
self._insert(
cache,
allocator,
shared_prefix + [5, 6],
session_id="session-a",
)
session_a_tail = self._stored_events(cache, StorageMedium.GPU)
self.assertEqual(len(session_a_tail), 1)
self.assertEqual(session_a_tail[0].parent_block_hash, shared_parent)
self.assertEqual(list(session_a_tail[0].token_ids), [5, 6])
self.assertEqual(session_a_tail[0].session_id, "session-a")
self._insert(
cache,
allocator,
shared_prefix + [7, 8],
session_id="session-b",
)
session_b_tail = self._stored_events(cache, StorageMedium.GPU)
self.assertEqual(len(session_b_tail), 1)
self.assertEqual(session_b_tail[0].parent_block_hash, shared_parent)
self.assertEqual(list(session_b_tail[0].token_ids), [7, 8])
self.assertEqual(session_b_tail[0].session_id, "session-b")
def test_session_id_and_cache_salt_are_both_attributed(self):
cache, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
cache.take_events()
self._insert(
cache,
allocator,
[1, 2, 3, 4],
cache_salt="tenant-a",
session_id="session-a",
)
stored = self._stored_events(cache, StorageMedium.GPU)
self.assertEqual(len(stored), 1)
self.assertEqual(stored[0].cache_salt, "tenant-a")
self.assertEqual(stored[0].session_id, "session-a")
def test_cache_salt_event_parentage_survives_node_split(self):
cache, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
cache.take_events()
@@ -1127,8 +1199,7 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
self._insert(cache, allocator, [1, 2, 5, 6], cache_salt="tenant-a")
branch = self._stored_events(cache, StorageMedium.GPU)
self.assertEqual(len(branch), 1)
self.assertIsInstance(branch[0], BlockStoredWithMetadata)
self.assertEqual(branch[0].metadata.cache_salt, "tenant-a")
self.assertEqual(branch[0].cache_salt, "tenant-a")
self.assertEqual(branch[0].parent_block_hash, original[0].block_hashes[0])
self.assertEqual(list(branch[0].token_ids), [5, 6])
@@ -1296,10 +1367,11 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
self.assertTrue(cache.tree_core.is_full_device_evicted(node))
self.assertTrue(cache.tree_core.is_backuped(node))
self._insert(cache, allocator, seq)
self._insert(cache, allocator, seq, session_id="session-a")
restored_gpu = self._stored_events(cache, StorageMedium.GPU)
self.assertFalse(cache.tree_core.is_full_device_evicted(node))
self.assertCountEqual(self._event_hashes(restored_gpu), stored_hashes)
self.assertEqual(restored_gpu[0].session_id, "session-a")
class UnifiedRadixCacheSuite: