fix(redis-semantic): reject unscoped scoped hits

This commit is contained in:
Ritwij Aryan Parmar 2026-05-29 23:45:43 -04:00
parent 63a491b1d7
commit a31147a0bf
2 changed files with 119 additions and 15 deletions

View file

@ -215,19 +215,25 @@ class RedisSemanticCache(BaseCache):
return metadata
def _get_scope_filter(self, **kwargs) -> Optional[Tuple[str, str]]:
def _get_scope_filters(self, **kwargs) -> List[Tuple[str, str]]:
metadata = self._get_metadata_from_kwargs(kwargs)
filters: List[Tuple[str, str]] = []
seen_fields: set[str] = set()
for field_name, metadata_key in self._SCOPE_FIELD_TO_METADATA_KEY:
value = metadata.get(metadata_key)
if value:
return field_name, str(value)
return None
if value and field_name not in seen_fields:
filters.append((field_name, str(value)))
seen_fields.add(field_name)
return filters
def _get_scope_filter(self, **kwargs) -> Optional[Tuple[str, str]]:
filters = self._get_scope_filters(**kwargs)
return filters[0] if filters else None
def _get_cache_filters(self, key: str, **kwargs) -> Dict[str, str]:
filters = {self.CACHE_KEY_FIELD_NAME: str(key)}
scope_filter = self._get_scope_filter(**kwargs)
if scope_filter is not None:
filters[scope_filter[0]] = scope_filter[1]
for field_name, value in self._get_scope_filters(**kwargs):
filters[field_name] = value
return filters
def _get_scope_filter_expression(self, **kwargs) -> Any:
@ -240,15 +246,34 @@ class RedisSemanticCache(BaseCache):
return Tag(scope_filter[0]) == scope_filter[1]
def _cache_hit_matches_scope(self, cache_hit: Dict[str, Any], **kwargs) -> bool:
scope_filter = self._get_scope_filter(**kwargs)
if scope_filter is None:
return True
scope_filters = self._get_scope_filters(**kwargs)
if not scope_filters:
return not any(
self._normalize_scope_value(cache_hit.get(field_name)) is not None
for field_name in self._scope_field_names()
)
field_name, expected_value = scope_filter
cached_value = cache_hit.get(field_name)
if isinstance(cached_value, bytes):
cached_value = cached_value.decode("utf-8")
return cached_value is not None and str(cached_value) == expected_value
for field_name, expected_value in scope_filters:
cached_value = self._normalize_scope_value(cache_hit.get(field_name))
if cached_value is not None and cached_value == expected_value:
return True
return False
@classmethod
def _scope_field_names(cls) -> Tuple[str, ...]:
return (
cls.API_KEY_HASH_FIELD_NAME,
cls.TEAM_ID_FIELD_NAME,
cls.USER_ID_FIELD_NAME,
)
@staticmethod
def _normalize_scope_value(value: Any) -> Optional[str]:
if value is None:
return None
if isinstance(value, bytes):
value = value.decode("utf-8")
return str(value)
def _get_ttl(self, **kwargs) -> Optional[int]:
"""

View file

@ -235,6 +235,51 @@ def test_redis_semantic_cache_set_cache_stores_cache_key_filter(monkeypatch):
)
def test_redis_semantic_cache_set_cache_stores_all_scope_filters(monkeypatch):
semantic_cache_mock = MagicMock()
custom_vectorizer_mock = MagicMock()
with patch.dict(
"sys.modules",
{
"redisvl.extensions.llmcache": MagicMock(SemanticCache=semantic_cache_mock),
"redisvl.utils.vectorize": MagicMock(
CustomTextVectorizer=custom_vectorizer_mock
),
},
):
from litellm.caching.redis_semantic_cache import RedisSemanticCache
monkeypatch.setenv("REDIS_HOST", "localhost")
monkeypatch.setenv("REDIS_PORT", "6379")
monkeypatch.setenv("REDIS_PASSWORD", "test_password")
redis_semantic_cache = RedisSemanticCache(similarity_threshold=0.8)
redis_semantic_cache.llmcache.store = MagicMock()
redis_semantic_cache.set_cache(
key="test_key",
value={"content": "Paris"},
messages=[{"content": "What is the capital of France?"}],
metadata={
"user_api_key_hash": "hashed-key",
"user_api_key_team_id": "team-123",
"user_id": "user-456",
},
)
redis_semantic_cache.llmcache.store.assert_called_once_with(
"What is the capital of France?",
"{'content': 'Paris'}",
filters={
RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key",
RedisSemanticCache.API_KEY_HASH_FIELD_NAME: "hashed-key",
RedisSemanticCache.TEAM_ID_FIELD_NAME: "team-123",
RedisSemanticCache.USER_ID_FIELD_NAME: "user-456",
},
)
def test_redis_semantic_cache_uses_isolated_index_for_old_schema(monkeypatch):
fallback_cache_mock = MagicMock()
semantic_cache_mock = MagicMock(
@ -395,6 +440,40 @@ def test_redis_semantic_cache_rejects_pre_scope_hit_for_scoped_request():
)
def test_redis_semantic_cache_rejects_scoped_hit_for_unscoped_request():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
cache_hit = {
"prompt": "What is the capital of France?",
"response": '{"content": "Paris"}',
"vector_distance": 0.1,
RedisSemanticCache.TEAM_ID_FIELD_NAME: "team-123",
}
assert not redis_semantic_cache._cache_hit_matches_scope(
cache_hit=cache_hit,
metadata={},
)
def test_redis_semantic_cache_matches_secondary_scope_on_multi_scope_hit():
from litellm.caching.redis_semantic_cache import RedisSemanticCache
redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache)
cache_hit = {
RedisSemanticCache.API_KEY_HASH_FIELD_NAME: "hashed-key",
RedisSemanticCache.TEAM_ID_FIELD_NAME: "team-123",
}
assert redis_semantic_cache._cache_hit_matches_scope(
cache_hit=cache_hit,
metadata={"user_api_key_team_id": "team-123"},
)
def test_redis_semantic_cache_builds_scope_filter_expression(monkeypatch):
class FakeTag:
def __init__(self, field_name):