From 26974ddd202155ec52bae518e233c1cff5ea7770 Mon Sep 17 00:00:00 2001 From: sid sri Date: Sun, 27 Sep 2026 23:52:36 +0530 Subject: [PATCH] test(caching): inject router in Valkey metadata tests --- .../caching/test_valkey_semantic_cache.py | 50 +++++++++++-------- 1 file changed, 30 insertions(+), 20 deletions(-) diff --git a/tests/test_litellm/caching/test_valkey_semantic_cache.py b/tests/test_litellm/caching/test_valkey_semantic_cache.py index 749658784ac..5b07a9efdbf 100644 --- a/tests/test_litellm/caching/test_valkey_semantic_cache.py +++ b/tests/test_litellm/caching/test_valkey_semantic_cache.py @@ -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