mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
fix(redis-semantic): reject unscoped scoped hits
This commit is contained in:
parent
63a491b1d7
commit
a31147a0bf
2 changed files with 119 additions and 15 deletions
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue