mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test(caching): inject router in Valkey metadata tests
This commit is contained in:
parent
feb0f4e63a
commit
26974ddd20
1 changed files with 30 additions and 20 deletions
|
|
@ -311,18 +311,19 @@ async def test_async_set_and_get_roundtrip():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_set_cache_passes_only_metadata_to_get_async_embedding():
|
||||
async def test_async_set_cache_passes_only_metadata_to_get_async_embedding(monkeypatch):
|
||||
import types
|
||||
|
||||
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
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]})
|
||||
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy.llm_router = router
|
||||
fake_proxy.llm_model_list = [{"model_name": cache.embedding_model}]
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
|
||||
|
||||
await cache.async_set_cache(
|
||||
key="cache-key",
|
||||
|
|
@ -333,23 +334,28 @@ async def test_async_set_cache_passes_only_metadata_to_get_async_embedding():
|
|||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert captured["metadata"] == {"user_api_key": "sk-test"}
|
||||
router.aembedding.assert_awaited_once()
|
||||
assert router.aembedding.call_args.kwargs["metadata"] == {
|
||||
"user_api_key": "sk-test",
|
||||
"semantic-cache-embedding": True,
|
||||
}
|
||||
async_client.hset.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_get_cache_passes_only_metadata_to_get_async_embedding():
|
||||
async def test_async_get_cache_passes_only_metadata_to_get_async_embedding(monkeypatch):
|
||||
import types
|
||||
|
||||
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
|
||||
router = MagicMock()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2, 0.3]}]})
|
||||
fake_proxy = types.ModuleType("litellm.proxy.proxy_server")
|
||||
fake_proxy.llm_router = router
|
||||
fake_proxy.llm_model_list = [{"model_name": cache.embedding_model}]
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy)
|
||||
|
||||
result = await cache.async_get_cache(
|
||||
key="cache-key",
|
||||
|
|
@ -360,7 +366,11 @@ async def test_async_get_cache_passes_only_metadata_to_get_async_embedding():
|
|||
)
|
||||
|
||||
assert result == {"content": "Paris"}
|
||||
assert captured["metadata"] == {"user_api_key": "sk-test"}
|
||||
router.aembedding.assert_awaited_once()
|
||||
assert router.aembedding.call_args.kwargs["metadata"] == {
|
||||
"user_api_key": "sk-test",
|
||||
"semantic-cache-embedding": True,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue