diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index 76b7f7d5b87..6baeb2d568f 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -240,7 +240,7 @@ class ValkeySemanticCache(RedisSemanticCache): print_verbose("No prompt provided for semantic caching") return - embedding = self._get_embedding(prompt) + embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata")) self._ensure_index_sync(len(embedding)) doc_key = self._doc_key(key) @@ -259,7 +259,7 @@ class ValkeySemanticCache(RedisSemanticCache): kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - embedding = self._get_embedding(prompt) + embedding = self._get_embedding(prompt, metadata=kwargs.get("metadata")) self._ensure_index_sync(len(embedding)) search_result = self.sync_client.ft(self.index_name).search( diff --git a/tests/test_litellm/caching/test_valkey_semantic_cache.py b/tests/test_litellm/caching/test_valkey_semantic_cache.py index d2df0a98e12..e56ed3ba4a2 100644 --- a/tests/test_litellm/caching/test_valkey_semantic_cache.py +++ b/tests/test_litellm/caching/test_valkey_semantic_cache.py @@ -353,6 +353,53 @@ async def test_async_get_cache_passes_only_metadata_to_get_async_embedding(): assert captured["metadata"] == {"user_api_key": "sk-test"} +def test_set_cache_passes_metadata_to_get_embedding(): + """Regression for #32324: the sync path must forward request metadata to the + embedding call so embedding spend is attributed, matching the async path and + RedisSemanticCache.""" + sync_client = MagicMock() + cache = _make_cache(sync_client=sync_client) + captured: dict[str, object] = {} + + def spy_embedding(prompt: str, metadata: dict | None = None) -> list[float]: + captured["metadata"] = metadata + return [0.1, 0.2, 0.3] + + cache._get_embedding = spy_embedding + + cache.set_cache( + key="cache-key", + value={"content": "Paris"}, + messages=[{"role": "user", "content": "What is the capital of France?"}], + metadata={"user_api_key": "sk-test"}, + ) + + assert captured["metadata"] == {"user_api_key": "sk-test"} + + +def test_get_cache_passes_metadata_to_get_embedding(): + """Regression for #32324: the sync get path must also forward metadata.""" + sync_client = MagicMock() + sync_client.ft.return_value.search.return_value = _search_result(0.1) + cache = _make_cache(sync_client=sync_client, similarity_threshold=0.8) + captured: dict[str, object] = {} + + def spy_embedding(prompt: str, metadata: dict | None = None) -> list[float]: + captured["metadata"] = dict(metadata) if metadata is not None else None + return [0.1, 0.2, 0.3] + + cache._get_embedding = spy_embedding + + result = cache.get_cache( + key="cache-key", + messages=[{"role": "user", "content": "capital of France?"}], + metadata={"user_api_key": "sk-test"}, + ) + + assert result == {"content": "Paris"} + assert captured["metadata"] == {"user_api_key": "sk-test"} + + @pytest.mark.asyncio async def test_async_get_cache_misses_below_threshold(): async_client = AsyncMock()