mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
test(caching): add regression tests for async_get_cache kwargs filtering
This commit is contained in:
parent
f870d7cee3
commit
1934d6e4cd
1 changed files with 58 additions and 0 deletions
|
|
@ -1,5 +1,8 @@
|
|||
import logging
|
||||
import re
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.types.caching import LiteLLMCacheType
|
||||
|
|
@ -146,3 +149,58 @@ def test_exact_cache_key_still_includes_prompt():
|
|||
model="gpt-4o-mini", messages=[{"role": "user", "content": "b"}]
|
||||
)
|
||||
assert key_a != key_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_cache_filters_kwargs_like_sync():
|
||||
"""async_get_cache must pass only messages/input/metadata to the backend,
|
||||
matching get_cache. Passing unfiltered kwargs (model, temperature, etc.)
|
||||
risks interference with prompt extraction or future backend changes."""
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
backend = MagicMock()
|
||||
backend.async_get_cache = AsyncMock(return_value=None)
|
||||
cache.cache = backend
|
||||
|
||||
messages = [{"role": "user", "content": "hello"}]
|
||||
metadata = {"user_api_key": "test-key"}
|
||||
|
||||
await cache.async_get_cache(
|
||||
model="gpt-4o-mini",
|
||||
messages=messages,
|
||||
metadata=metadata,
|
||||
temperature=0.7,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
backend.async_get_cache.assert_called_once()
|
||||
_, call_kwargs = backend.async_get_cache.call_args
|
||||
assert "model" not in call_kwargs
|
||||
assert "temperature" not in call_kwargs
|
||||
assert "stream" not in call_kwargs
|
||||
assert call_kwargs["messages"] is messages
|
||||
assert call_kwargs["metadata"] == metadata
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_cache_propagates_semantic_similarity_metadata():
|
||||
"""async_get_cache must copy semantic-similarity from the backend's metadata
|
||||
copy back into the caller's original metadata dict."""
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
|
||||
async def fake_get(key, **kwargs):
|
||||
kwargs.get("metadata", {})["semantic-similarity"] = 0.95
|
||||
return None
|
||||
|
||||
backend = MagicMock()
|
||||
backend.async_get_cache = AsyncMock(side_effect=fake_get)
|
||||
cache.cache = backend
|
||||
|
||||
original_metadata: dict = {"user_api_key": "k"}
|
||||
|
||||
await cache.async_get_cache(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
metadata=original_metadata,
|
||||
)
|
||||
|
||||
assert original_metadata.get("semantic-similarity") == 0.95
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue