mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix: filter sensitive cache lookup kwargs
This commit is contained in:
parent
4da8111faf
commit
37406e7948
2 changed files with 106 additions and 4 deletions
|
|
@ -497,6 +497,34 @@ class Cache:
|
|||
return cached_response
|
||||
return cached_result
|
||||
|
||||
@staticmethod
|
||||
def _get_safe_cache_lookup_kwargs(kwargs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
cache_lookup_kwargs: Dict[str, Any] = {}
|
||||
for prompt_kwarg in ("messages", "input"):
|
||||
if prompt_kwarg in kwargs:
|
||||
cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg]
|
||||
|
||||
if isinstance(kwargs.get("metadata"), dict):
|
||||
cache_lookup_kwargs["metadata"] = {}
|
||||
|
||||
return cache_lookup_kwargs
|
||||
|
||||
@staticmethod
|
||||
def _update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs: Dict[str, Any], cache_lookup_kwargs: Dict[str, Any]
|
||||
) -> None:
|
||||
original_metadata = original_kwargs.get("metadata")
|
||||
cache_lookup_metadata = cache_lookup_kwargs.get("metadata")
|
||||
if not isinstance(original_metadata, dict) or not isinstance(
|
||||
cache_lookup_metadata, dict
|
||||
):
|
||||
return
|
||||
|
||||
if "semantic-similarity" in cache_lookup_metadata:
|
||||
original_metadata["semantic-similarity"] = cache_lookup_metadata[
|
||||
"semantic-similarity"
|
||||
]
|
||||
|
||||
def get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs):
|
||||
"""
|
||||
Retrieves the cached result for the given arguments.
|
||||
|
|
@ -522,10 +550,19 @@ class Cache:
|
|||
or cache_control_args.get("s-max-age")
|
||||
or float("inf")
|
||||
)
|
||||
cache_lookup_kwargs = self._get_safe_cache_lookup_kwargs(kwargs)
|
||||
if dynamic_cache_object is not None:
|
||||
cached_result = dynamic_cache_object.get_cache(cache_key, **kwargs)
|
||||
cached_result = dynamic_cache_object.get_cache(
|
||||
cache_key, **cache_lookup_kwargs
|
||||
)
|
||||
else:
|
||||
cached_result = self.cache.get_cache(cache_key, **kwargs)
|
||||
cached_result = self.cache.get_cache(
|
||||
cache_key, **cache_lookup_kwargs
|
||||
)
|
||||
self._update_metadata_from_cache_lookup_kwargs(
|
||||
original_kwargs=kwargs,
|
||||
cache_lookup_kwargs=cache_lookup_kwargs,
|
||||
)
|
||||
return self._get_cache_logic(
|
||||
cached_result=cached_result, max_age=max_age
|
||||
)
|
||||
|
|
|
|||
|
|
@ -890,7 +890,73 @@ def test_cache_get_cache_passes_responses_input_to_backend_cache():
|
|||
"test_key",
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
cache={},
|
||||
)
|
||||
|
||||
|
||||
def test_cache_get_cache_filters_sensitive_kwargs_from_backend_cache():
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
cache = Cache.__new__(Cache)
|
||||
cache.cache = MagicMock()
|
||||
cache.should_use_cache = MagicMock(return_value=True)
|
||||
cache.get_cache_key = MagicMock(return_value="test_key")
|
||||
cache._get_cache_logic = MagicMock(return_value={"content": "Paris"})
|
||||
|
||||
def _cache_hit(_cache_key, **cache_kwargs):
|
||||
cache_kwargs["metadata"]["semantic-similarity"] = 0.7
|
||||
return {"content": "Paris"}
|
||||
|
||||
cache.cache.get_cache = MagicMock(side_effect=_cache_hit)
|
||||
|
||||
metadata = {"user_api_key": "sk-secret", "trace_id": "trace-id"}
|
||||
result = cache.get_cache(
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
cache={"s-maxage": 10},
|
||||
api_key="sk-secret",
|
||||
headers={"authorization": "Bearer sk-secret"},
|
||||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
assert metadata == {
|
||||
"user_api_key": "sk-secret",
|
||||
"trace_id": "trace-id",
|
||||
"semantic-similarity": 0.7,
|
||||
}
|
||||
|
||||
forwarded_kwargs = cache.cache.get_cache.call_args.kwargs
|
||||
assert forwarded_kwargs == {
|
||||
"input": "What is the capital of France?",
|
||||
"metadata": {"semantic-similarity": 0.7},
|
||||
}
|
||||
assert forwarded_kwargs["metadata"] is not metadata
|
||||
cache._get_cache_logic.assert_called_once_with(
|
||||
cached_result={"content": "Paris"},
|
||||
max_age=10,
|
||||
)
|
||||
|
||||
|
||||
def test_cache_get_cache_filters_sensitive_kwargs_without_metadata():
|
||||
from litellm.caching.caching import Cache
|
||||
|
||||
cache = Cache.__new__(Cache)
|
||||
cache.cache = MagicMock()
|
||||
cache.cache.get_cache = MagicMock(return_value={"content": "Paris"})
|
||||
cache.should_use_cache = MagicMock(return_value=True)
|
||||
cache.get_cache_key = MagicMock(return_value="test_key")
|
||||
cache._get_cache_logic = MagicMock(return_value={"content": "Paris"})
|
||||
|
||||
result = cache.get_cache(
|
||||
input="What is the capital of France?",
|
||||
cache={"s-maxage": 10},
|
||||
api_key="sk-secret",
|
||||
headers={"authorization": "Bearer sk-secret"},
|
||||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
cache.cache.get_cache.assert_called_once_with(
|
||||
"test_key",
|
||||
input="What is the capital of France?",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -917,7 +983,6 @@ def test_cache_get_cache_passes_responses_input_to_dynamic_cache():
|
|||
"test_key",
|
||||
input="What is the capital of France?",
|
||||
metadata=metadata,
|
||||
cache={},
|
||||
)
|
||||
cache._get_cache_logic.assert_called_once_with(
|
||||
cached_result={"content": "Paris"},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue