From ed66ee312cba4c5e07cd51da9b80842a5c8abc76 Mon Sep 17 00:00:00 2001 From: Shivam Rawat Date: Mon, 6 Jul 2026 22:52:03 -0700 Subject: [PATCH] fix(caching): pass only metadata to valkey semantic async embedding (#32295) * fix(caching): pass only metadata to valkey semantic async embedding ValkeySemanticCache async get/set passed **kwargs into _get_async_embedding, which raised TypeError on cache_key and other fields and silently skipped all cache writes. Match redis-semantic by forwarding metadata only. Co-authored-by: Cursor * test(caching): add async_get_cache embedding call regression test Mirror the async_set_cache spy test so async_get_cache passing **kwargs into _get_async_embedding is caught by a real signature, not AsyncMock. Co-authored-by: Cursor --------- Co-authored-by: Shivam Rawat Co-authored-by: Cursor --- litellm/caching/valkey_semantic_cache.py | 4 +- .../caching/test_valkey_semantic_cache.py | 53 +++++++++++++++++++ 2 files changed, 55 insertions(+), 2 deletions(-) diff --git a/litellm/caching/valkey_semantic_cache.py b/litellm/caching/valkey_semantic_cache.py index 746e91207d8..76b7f7d5b87 100644 --- a/litellm/caching/valkey_semantic_cache.py +++ b/litellm/caching/valkey_semantic_cache.py @@ -279,7 +279,7 @@ class ValkeySemanticCache(RedisSemanticCache): print_verbose("No prompt provided for semantic caching") return - embedding = await self._get_async_embedding(prompt, **kwargs) + embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) await self._ensure_index_async(len(embedding)) doc_key = self._doc_key(key) @@ -298,7 +298,7 @@ class ValkeySemanticCache(RedisSemanticCache): kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - embedding = await self._get_async_embedding(prompt, **kwargs) + embedding = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata")) await self._ensure_index_async(len(embedding)) search_result = await self.async_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 44b9f061998..d2df0a98e12 100644 --- a/tests/test_litellm/caching/test_valkey_semantic_cache.py +++ b/tests/test_litellm/caching/test_valkey_semantic_cache.py @@ -300,6 +300,59 @@ async def test_async_set_and_get_roundtrip(): assert metadata["semantic-similarity"] == pytest.approx(0.95) +@pytest.mark.asyncio +async def test_async_set_cache_passes_only_metadata_to_get_async_embedding(): + async_client = AsyncMock() + async_client.ft = _async_ft(0.05) + cache = _make_cache(async_client=async_client) + captured: dict[str, object] = {} + + async def spy_embedding(prompt: str, metadata: dict | None = None) -> list[float]: + captured["prompt"] = prompt + captured["metadata"] = metadata + return [0.1, 0.2, 0.3] + + cache._get_async_embedding = spy_embedding + + await cache.async_set_cache( + key="cache-key", + value={"content": "Paris"}, + messages=[{"role": "user", "content": "What is the capital of France?"}], + metadata={"user_api_key": "sk-test"}, + cache_key="abc123", + custom_llm_provider="openai", + ) + + assert captured["metadata"] == {"user_api_key": "sk-test"} + async_client.hset.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_async_get_cache_passes_only_metadata_to_get_async_embedding(): + async_client = AsyncMock() + async_client.ft = _async_ft(0.05) + cache = _make_cache(async_client=async_client) + captured: dict[str, object] = {} + + async def spy_embedding(prompt: str, metadata: dict | None = None) -> list[float]: + captured["prompt"] = prompt + captured["metadata"] = dict(metadata) if metadata is not None else None + return [0.1, 0.2, 0.3] + + cache._get_async_embedding = spy_embedding + + result = await cache.async_get_cache( + key="cache-key", + messages=[{"role": "user", "content": "What is the capital of France?"}], + metadata={"user_api_key": "sk-test"}, + cache_key="abc123", + custom_llm_provider="openai", + ) + + 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()