feat: add cache salt support to KV cache events (#30827)

Signed-off-by: jthomson04 <jwillthomson19@gmail.com>
This commit is contained in:
jthomson04
2026-08-12 16:14:04 -07:00
committed by GitHub
parent 198b7e9240
commit 385903b0ac
43 changed files with 753 additions and 70 deletions
@@ -5,6 +5,7 @@ import sys
import types
import unittest
from array import array
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from sglang.srt.mem_cache.evict_policy import (
@@ -17,6 +18,7 @@ from sglang.srt.mem_cache.evict_policy import (
SLRUStrategy,
)
from sglang.srt.mem_cache.utils import (
compute_node_event_hash_values,
compute_node_hash_values,
get_eviction_strategy,
get_hash_str,
@@ -54,9 +56,10 @@ def _legacy_page_hashes(key, page_size, prior_hash=None):
class _HashKey:
def __init__(self, token_ids, is_bigram=False):
def __init__(self, token_ids, is_bigram=False, cache_salt=None):
self.token_ids = token_ids
self.is_bigram = is_bigram
self.cache_salt = cache_salt
def __len__(self):
if self.is_bigram:
@@ -68,8 +71,12 @@ class _HashKey:
start = index.start or 0
stop = index.stop if index.stop is not None else len(self)
if self.is_bigram:
return _HashKey(self.token_ids[start : stop + 1], is_bigram=True)
return _HashKey(self.token_ids[start:stop])
return _HashKey(
self.token_ids[start : stop + 1],
is_bigram=True,
cache_salt=self.cache_salt,
)
return _HashKey(self.token_ids[start:stop], cache_salt=self.cache_salt)
if self.is_bigram:
return (self.token_ids[index], self.token_ids[index + 1])
return self.token_ids[index]
@@ -298,6 +305,7 @@ class TestComputeNodeHashValues(unittest.TestCase):
node = MagicMock()
node.key = key
node.parent = parent
node.event_hash_value = None
if parent is not None:
parent.hash_value = parent_hash_values
return node
@@ -318,6 +326,60 @@ class TestComputeNodeHashValues(unittest.TestCase):
_legacy_page_hashes(key, page_size=page_size),
)
def test_cache_salt_seeds_root_hash_chain(self):
key = _HashKey(array("q", range(1, 17)), cache_salt="tenant-a")
seed = hashlib.sha256(b"sglang-cache-salt-v1\0tenant-a").hexdigest()
self.assertEqual(
compute_node_event_hash_values(self._make_node(key), page_size=8),
_legacy_page_hashes(key, page_size=8, prior_hash=seed),
)
self.assertEqual(
compute_node_hash_values(self._make_node(key), page_size=8),
_legacy_page_hashes(key, page_size=8),
)
other = _HashKey(array("q", range(1, 17)), cache_salt="tenant-b")
self.assertNotEqual(
compute_node_event_hash_values(self._make_node(key), page_size=8),
compute_node_event_hash_values(self._make_node(other), page_size=8),
)
def test_cache_salt_event_hashes_are_memoized(self):
node = self._make_node(
_HashKey(array("q", range(1, 17)), cache_salt="tenant-a")
)
with patch(
"sglang.srt.mem_cache.utils.get_hash_str", wraps=get_hash_str
) as mock_get_hash_str:
first = compute_node_event_hash_values(node, page_size=8)
second = compute_node_event_hash_values(node, page_size=8)
self.assertIs(first, second)
mock_get_hash_str.assert_called_once()
def test_cache_salt_event_hash_walk_is_iterative(self):
root = SimpleNamespace(
key=_HashKey(array("q")),
parent=None,
hash_value=[],
event_hash_value=None,
)
node = root
path = []
for token_id in range(1, 1102):
node = SimpleNamespace(
key=_HashKey(array("q", [token_id]), cache_salt="tenant-a"),
parent=node,
hash_value=None,
event_hash_value=None,
)
path.append(node)
result = compute_node_event_hash_values(node, page_size=1)
self.assertEqual(len(result), 1)
self.assertTrue(all(item.event_hash_value is not None for item in path))
def test_parent_hash_is_used_only_when_parent_has_nonempty_key_and_hash(self):
parent = MagicMock()
parent.key = _HashKey(array("q", range(1, 17)))