From 1260f0967937ab70fffa48c67c8ec9a1838dc29f Mon Sep 17 00:00:00 2001 From: Kent Date: Mon, 29 Jun 2026 00:06:19 +0800 Subject: [PATCH] fix(cache): resolve semantic-cache embedding model via Router model-group lookup resolve_embedding_router gated on exact model-name membership, so wildcard/pattern (bedrock/*) and visible model_group_alias embedding models fell back to a direct litellm.embedding() call and lost deployment-level auth (e.g. Bedrock aws_role_name), the same class of failure as #28244. Resolve via Router.get_model_list(model_name=...), which unifies exact + visible alias + provider-prefixed wildcard/pattern, and drop the now-redundant llm_model_list arg. get_model_list is a listing resolver, not a faithful 'would the Router route this' predicate: hidden model_group_alias entries, a bare model name matched only against a provider-prefixed wildcard, a raw litellm_params.model / model_info.id deployment address, and team-public names without a team_id still resolve empty and keep falling back to a direct embedding call. These are niche embedding-model configs and degrade to current behavior; documented as known gaps. --- litellm/caching/_embedding_router.py | 10 +-- litellm/caching/qdrant_semantic_cache.py | 10 +-- litellm/caching/redis_semantic_cache.py | 10 +-- .../caching/test_embedding_router.py | 75 +++++++++++++------ .../caching/test_qdrant_semantic_cache.py | 21 ++---- .../caching/test_redis_semantic_cache.py | 5 +- 6 files changed, 72 insertions(+), 59 deletions(-) diff --git a/litellm/caching/_embedding_router.py b/litellm/caching/_embedding_router.py index ec886b14020..0210dee6d60 100644 --- a/litellm/caching/_embedding_router.py +++ b/litellm/caching/_embedding_router.py @@ -5,8 +5,8 @@ configured embedding model is a proxy Router deployment, embeddings must run through the Router so per-deployment auth (e.g. Bedrock aws_role_name) is applied. Otherwise fall back to a direct litellm embedding call. -This module is dependency-injected: callers pass the proxy ``llm_router`` and -``llm_model_list`` in, so the decision logic is unit-testable without importing +This module is dependency-injected: callers pass the proxy ``llm_router`` in, so +the decision logic is unit-testable without importing ``litellm.proxy.proxy_server``. """ @@ -21,15 +21,11 @@ if TYPE_CHECKING: def resolve_embedding_router( embedding_model: str, llm_router: Router | None, - llm_model_list: list[dict[str, Any]] | None, ) -> Router | None: """Return ``llm_router`` iff it serves ``embedding_model`` as a deployment.""" if llm_router is None: return None - router_model_names: list[str] = ( - [m["model_name"] for m in llm_model_list if "model_name" in m] if llm_model_list is not None else [] - ) - if embedding_model in router_model_names: + if llm_router.get_model_list(model_name=embedding_model): return llm_router return None diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index 5ed1bb47eba..885d7426b97 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -191,12 +191,11 @@ class QdrantSemanticCache(BaseCache): def _get_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> EmbeddingResponse: """Embed via the proxy Router when it serves the model, else direct.""" try: - from litellm.proxy.proxy_server import llm_model_list, llm_router + from litellm.proxy.proxy_server import llm_router except ImportError: - llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router = resolve_embedding_router(self.embedding_model, llm_router) if router is not None: return router.embedding( model=self.embedding_model, @@ -212,12 +211,11 @@ class QdrantSemanticCache(BaseCache): async def _get_async_embedding(self, prompt: str, metadata: Dict[str, Any] | None = None) -> EmbeddingResponse: try: - from litellm.proxy.proxy_server import llm_model_list, llm_router + from litellm.proxy.proxy_server import llm_router except ImportError: - llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router = resolve_embedding_router(self.embedding_model, llm_router) if router is not None: return await router.aembedding( model=self.embedding_model, diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index d4288cc777c..ca6e30d1ef3 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -313,12 +313,11 @@ class RedisSemanticCache(BaseCache): mirroring ``_get_async_embedding``; otherwise embeds directly. """ try: - from litellm.proxy.proxy_server import llm_model_list, llm_router + from litellm.proxy.proxy_server import llm_router except ImportError: - llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router = resolve_embedding_router(self.embedding_model, llm_router) if router is not None: embedding_response = cast( EmbeddingResponse, @@ -483,12 +482,11 @@ class RedisSemanticCache(BaseCache): List[float]: The embedding vector """ try: - from litellm.proxy.proxy_server import llm_model_list, llm_router + from litellm.proxy.proxy_server import llm_router except ImportError: - llm_model_list = None llm_router = None - router = resolve_embedding_router(self.embedding_model, llm_router, llm_model_list) + router = resolve_embedding_router(self.embedding_model, llm_router) try: if router is not None: embedding_response = await router.aembedding( diff --git a/tests/test_litellm/caching/test_embedding_router.py b/tests/test_litellm/caching/test_embedding_router.py index 550095a112a..2eeb773d28a 100644 --- a/tests/test_litellm/caching/test_embedding_router.py +++ b/tests/test_litellm/caching/test_embedding_router.py @@ -4,47 +4,76 @@ from unittest.mock import MagicMock sys.path.insert(0, os.path.abspath("../../..")) +import litellm + from litellm.caching._embedding_router import ( build_router_embedding_metadata, resolve_embedding_router, ) -def test_resolve_returns_router_when_model_is_a_deployment(): - router = MagicMock() +def test_resolve_routes_exact_name_model_via_real_router(): + router = litellm.Router( + model_list=[ + { + "model_name": "sem-embed", + "litellm_params": {"model": "text-embedding-3-small"}, + } + ] + ) + assert resolve_embedding_router("sem-embed", router) is router + + +def test_resolve_routes_provider_prefixed_wildcard_via_real_router(): + router = litellm.Router( + model_list=[ + {"model_name": "bedrock/*", "litellm_params": {"model": "bedrock/*"}} + ] + ) + # bedrock/amazon.titan-embed-text-v2:0 is NOT an exact model_name; only the + # bedrock/* pattern serves it. The old exact-name code returned None here. assert ( - resolve_embedding_router("sem-embed", router, [{"model_name": "sem-embed"}]) + resolve_embedding_router("bedrock/amazon.titan-embed-text-v2:0", router) is router ) -def test_resolve_returns_none_when_model_not_in_router(): - router = MagicMock() - assert ( - resolve_embedding_router("sem-embed", router, [{"model_name": "other"}]) is None +def test_resolve_routes_visible_model_group_alias_via_real_router(): + router = litellm.Router( + model_list=[ + { + "model_name": "real-embed", + "litellm_params": {"model": "text-embedding-3-small"}, + } + ], + model_group_alias={"aliased-embed": "real-embed"}, ) + # aliased-embed is only reachable through model_group_alias. + assert resolve_embedding_router("aliased-embed", router) is router + + +def test_resolve_returns_none_when_real_router_does_not_serve_model(): + router = litellm.Router( + model_list=[ + { + "model_name": "other-embed", + "litellm_params": {"model": "text-embedding-3-small"}, + } + ] + ) + assert resolve_embedding_router("sem-embed", router) is None def test_resolve_returns_none_when_router_is_none(): - assert ( - resolve_embedding_router("sem-embed", None, [{"model_name": "sem-embed"}]) - is None - ) + assert resolve_embedding_router("sem-embed", None) is None -def test_resolve_returns_none_when_model_list_is_none(): +def test_resolve_returns_none_when_get_model_list_returns_none(): + # get_model_list is annotated Optional[List]; in practice it returns [], + # but pin the falsy-None path so the `if ...:` gate stays correct. router = MagicMock() - assert resolve_embedding_router("sem-embed", router, None) is None - - -def test_resolve_skips_entries_missing_model_name(): - router = MagicMock() - model_list = [ - {"litellm_params": {"model": "bedrock/x"}}, - {"model_name": "sem-embed"}, - ] - assert resolve_embedding_router("sem-embed", router, model_list) is router - assert resolve_embedding_router("other", router, [{"litellm_params": {}}]) is None + router.get_model_list = MagicMock(return_value=None) + assert resolve_embedding_router("sem-embed", router) is None def test_build_metadata_preserves_request_fields_and_adds_flag(): diff --git a/tests/test_litellm/caching/test_qdrant_semantic_cache.py b/tests/test_litellm/caching/test_qdrant_semantic_cache.py index 67d4e2d9892..a1e2a80b282 100644 --- a/tests/test_litellm/caching/test_qdrant_semantic_cache.py +++ b/tests/test_litellm/caching/test_qdrant_semantic_cache.py @@ -22,7 +22,6 @@ def test_qdrant_semantic_cache_initialization(monkeypatch): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -82,7 +81,6 @@ def test_qdrant_semantic_cache_get_cache_hit(): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -162,7 +160,6 @@ def test_qdrant_semantic_cache_rejects_unscoped_cache_hit(): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - mock_response = MagicMock() mock_response.status_code = 200 mock_response.json.return_value = {"result": {"exists": True}} @@ -323,7 +320,6 @@ def test_qdrant_semantic_cache_get_cache_miss(): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -379,7 +375,6 @@ async def test_qdrant_semantic_cache_async_get_cache_hit(): "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" ) as mock_async_client, ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -470,7 +465,6 @@ async def test_qdrant_semantic_cache_async_get_cache_miss(): "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" ) as mock_async_client, ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -530,7 +524,6 @@ def test_qdrant_semantic_cache_set_cache(): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -596,7 +589,6 @@ async def test_qdrant_semantic_cache_async_set_cache(): "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" ) as mock_async_client, ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -666,7 +658,6 @@ def test_qdrant_semantic_cache_custom_vector_size(): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection does NOT exist (so it will be created) mock_exists_response = MagicMock() mock_exists_response.status_code = 200 @@ -727,7 +718,6 @@ def test_qdrant_semantic_cache_default_vector_size(): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -763,7 +753,6 @@ def test_qdrant_semantic_cache_large_vector_size(): ) as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection does NOT exist (so it will be created) mock_exists_response = MagicMock() mock_exists_response.status_code = 200 @@ -809,10 +798,10 @@ def test_qdrant_semantic_cache_large_vector_size(): assert create_payload["vectors"]["size"] == 4096 -def _router_proxy_module(router, model_name): +def _router_proxy_module(router): mod = types.ModuleType("litellm.proxy.proxy_server") mod.llm_router = router - mod.llm_model_list = [{"model_name": model_name}] + mod.llm_model_list = None return mod @@ -832,13 +821,14 @@ def test_qdrant_sync_get_cache_routes_through_router(monkeypatch): cache.sync_client.post.return_value = search_response router = MagicMock() + router.get_model_list = MagicMock(return_value=[{"model_name": "sem-embed"}]) router.embedding = MagicMock( return_value={"data": [{"embedding": [0.3, 0.3, 0.3]}]} ) monkeypatch.setitem( sys.modules, "litellm.proxy.proxy_server", - _router_proxy_module(router, "sem-embed"), + _router_proxy_module(router), ) with patch("litellm.embedding") as direct_embed: @@ -892,11 +882,12 @@ async def test_qdrant_async_embedding_forwards_full_metadata(monkeypatch): cache.embedding_model = "sem-embed" router = MagicMock() + router.get_model_list = MagicMock(return_value=[{"model_name": "sem-embed"}]) router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]}) monkeypatch.setitem( sys.modules, "litellm.proxy.proxy_server", - _router_proxy_module(router, "sem-embed"), + _router_proxy_module(router), ) await cache._get_async_embedding( diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py index 1d3129d6467..9bd8bdb5f0e 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -901,10 +901,11 @@ def test_redis_get_embedding_routes_through_router(monkeypatch): cache.embedding_model = "sem-embed" router = MagicMock() + router.get_model_list = MagicMock(return_value=[{"model_name": "sem-embed"}]) router.embedding = MagicMock(return_value={"data": [{"embedding": [0.5, 0.6]}]}) fake_proxy = types.ModuleType("litellm.proxy.proxy_server") fake_proxy.llm_router = router - fake_proxy.llm_model_list = [{"model_name": "sem-embed"}] + fake_proxy.llm_model_list = None monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy) with patch("litellm.embedding") as direct_embed: @@ -1145,10 +1146,10 @@ async def test_redis_async_embedding_forwards_full_metadata(monkeypatch): cache.embedding_model = "sem-embed" router = MagicMock() + router.get_model_list = MagicMock(return_value=[{"model_name": "sem-embed"}]) router.aembedding = AsyncMock(return_value={"data": [{"embedding": [0.1, 0.2]}]}) fake_proxy = types.ModuleType("litellm.proxy.proxy_server") fake_proxy.llm_router = router - fake_proxy.llm_model_list = [{"model_name": "sem-embed"}] monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_proxy) await cache._get_async_embedding(