mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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 <cursoragent@cursor.com> * 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 <cursoragent@cursor.com> --------- Co-authored-by: Shivam Rawat <shivamrawat@Shivams-MacBook-Pro.local> Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
7d6a080d3f
commit
ed66ee312c
2 changed files with 55 additions and 2 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue