From a31147a0bf160d77c4d4d228017ac693574857f4 Mon Sep 17 00:00:00 2001 From: Ritwij Aryan Parmar Date: Fri, 29 May 2026 23:45:43 -0400 Subject: [PATCH] fix(redis-semantic): reject unscoped scoped hits --- litellm/caching/redis_semantic_cache.py | 55 +++++++++---- .../caching/test_redis_semantic_cache.py | 79 +++++++++++++++++++ 2 files changed, 119 insertions(+), 15 deletions(-) diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index 7f651581c94..53d94fe7eac 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -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]: """ diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py index df82895d851..16e68ab70f5 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -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):