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.
This commit is contained in:
Kent 2026-06-29 00:06:19 +08:00
parent 8e30cfbeb1
commit 1260f09679
6 changed files with 72 additions and 59 deletions

View file

@ -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

View file

@ -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,

View file

@ -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(

View file

@ -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():

View file

@ -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(

View file

@ -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(