diff --git a/tests/router_unit_tests/test_router_embedding_integration.py b/tests/router_unit_tests/test_router_embedding_integration.py index e10ba0f0962..b295a351985 100644 --- a/tests/router_unit_tests/test_router_embedding_integration.py +++ b/tests/router_unit_tests/test_router_embedding_integration.py @@ -6,6 +6,7 @@ need to be properly propagated through the router to the LLM API. """ import json +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -13,6 +14,7 @@ import pytest import respx import litellm +import litellm.router_strategy.simple_shuffle as simple_shuffle from litellm import Router from litellm.llms.base_llm.vector_store.transformation import ( LiteLLMVectorStoreEmbeddingExecutor, @@ -383,6 +385,55 @@ class TestRouterEmbeddingIntegration: # The call should succeed mock_aembedding.assert_called_once() + def test_sync_embedding_respects_model_access_group_scoping(self, monkeypatch): + """ + Regression test for #31260: sync Router._embedding must forward the + caller's request_kwargs to deployment selection. Without the forward, + _filter_deployments_by_model_access_groups short-circuits + (request_kwargs is None) and every deployment in the model group is + treated as globally accessible. random.choice is forced to the last + candidate so an unfiltered deployment list fails this test + deterministically instead of on 50% of runs. + """ + model_list = [ + { + "model_name": "grouped-embed", + "litellm_params": { + "model": "text-embedding-3-small", + "api_key": "eng-key", + }, + "model_info": {"access_groups": ["engineering"]}, + }, + { + "model_name": "grouped-embed", + "litellm_params": { + "model": "text-embedding-3-small", + "api_key": "fin-key", + }, + "model_info": {"access_groups": ["finance"]}, + }, + ] + + router = Router(model_list=model_list) + engineering_auth = SimpleNamespace(models=["engineering"], team_models=[]) + + fake_random = MagicMock() + fake_random.choice.side_effect = lambda seq: seq[-1] + + monkeypatch.setattr(litellm, "embedding", MagicMock(return_value=MagicMock(data=[{"embedding": [0.1, 0.2]}]))) + monkeypatch.setattr(simple_shuffle, "random", fake_random) + + mock_embedding = litellm.embedding + router.embedding( + model="grouped-embed", + input=["hello"], + metadata={"user_api_key_auth": engineering_auth}, + ) + + # The finance deployment must have been filtered out; only the + # engineering deployment can be selected. + assert mock_embedding.call_args[1]["api_key"] == "eng-key" + def test_embedding_with_timeout_from_router(self): """ Test that timeout settings from router config are propagated. diff --git a/tests/test_litellm/caching/test_caching.py b/tests/test_litellm/caching/test_caching.py index 4d0ec0fb677..7fb2f9a75f2 100644 --- a/tests/test_litellm/caching/test_caching.py +++ b/tests/test_litellm/caching/test_caching.py @@ -1,5 +1,6 @@ import logging import re +from unittest.mock import MagicMock import pytest @@ -252,3 +253,33 @@ def test_exact_cache_key_includes_anthropic_messages_params(anthropic_param): assert baseline != cache.get_cache_key( model="claude-sonnet-4-5", messages=messages, **anthropic_param ) + + +def test_sync_get_cache_forwards_metadata_subset_to_backend(): + """ + Regression test for #31260: the sync cache read path used to flatten + metadata to {}, so semantic-cache backends never saw the caller's auth + context (team id / user_api_key_auth) on a get, while the write path and + the async read path forwarded full metadata. The metadata subset must be + preserved so sync semantic-cache lookups can scope and authenticate the + same way async ones do. + """ + cache = Cache(type=LiteLLMCacheType.LOCAL) + backend = MagicMock() + backend.get_cache.return_value = None + + cache.get_cache( + dynamic_cache_object=backend, + cache={"use-cache": True}, + model="text-embedding-3-small", + input=["hello"], + metadata={ + "user_api_key_team_id": "team-a", + "user_api_key": "k-abc", + }, + ) + + assert backend.get_cache.called + sent_metadata = backend.get_cache.call_args.kwargs["metadata"] + assert sent_metadata["user_api_key_team_id"] == "team-a" + assert sent_metadata["user_api_key"] == "k-abc"