fix: filter sensitive cache lookup kwargs

This commit is contained in:
Dávid Balatoni 2026-05-23 22:50:59 +02:00
parent 4da8111faf
commit 37406e7948
No known key found for this signature in database
2 changed files with 106 additions and 4 deletions

View file

@ -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
)

View file

@ -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"},