diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index cb688860280..bc3bcd233f0 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -43,6 +43,7 @@ from litellm.litellm_core_utils.llm_response_utils.get_headers import ( ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.auth_utils import check_response_size_is_safe from litellm.proxy.common_utils.callback_utils import ( get_logging_caching_headers, @@ -53,6 +54,7 @@ from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict +from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.types.guardrails import GuardrailEventHooks from litellm.types.router import RouterRateLimitError from litellm.types.utils import ServerToolUse @@ -384,6 +386,39 @@ async def _authorize_response_file_search_vector_stores( ) +async def _resolve_per_request_model_group_alias( + requested_model: object, + router_settings: Mapping[str, object], + user_api_key_dict: UserAPIKeyAuth, + llm_router: Router, +) -> str | None: + """ + Resolve ``router_settings.model_group_alias`` coming from a key or team. + + The Router only ever resolves aliases from its own instance attribute, which + holds the global config map and is shared across requests, so a per-request + map has to be applied here instead of being forwarded to the Router. + + Model access was authorized against the requested group, so the target is + authorized in its own right before the rewrite; a key that may not call the + target gets the usual 403 rather than being quietly served it. + + Returns the target model group, or None when no alias applies. + """ + if not isinstance(requested_model, str): + return None + target = resolve_model_group_alias(router_settings.get("model_group_alias"), requested_model) + if target is None or target == requested_model: + return None + await can_key_call_resolved_model( + model=target, + llm_model_list=llm_router.model_list, + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + return target + + async def _parse_event_data_for_error(event_line: str | bytes) -> int | None: """Parses an event line and returns an error code if present, else None.""" event_line = event_line.decode("utf-8") if isinstance(event_line, bytes) else event_line @@ -1285,6 +1320,35 @@ class ProxyBaseLLMRequestProcessing: ): self.data["model"] = user_api_key_dict.aliases[self.data["model"]] + # Apply hierarchical router_settings (Key > Team) + # Global router_settings are already on the Router object itself. + # This sits with the other alias rewrites, and ahead of the guardrail + # merge and the pre-call hooks, so everything that keys off the model + # group -- model-level guardrails, per-model budgets and rate limits, + # the logging object -- sees the group that will actually serve. + if llm_router is not None and proxy_config is not None: + from litellm.proxy.proxy_server import prisma_client + + router_settings = await proxy_config._get_hierarchical_router_settings( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + ) + + # If router_settings found (from key or team), apply them + # Pass settings as per-request overrides instead of creating a new Router + # This avoids expensive Router instantiation on each request + if router_settings is not None: + self.data["router_settings_override"] = router_settings + alias_target = await _resolve_per_request_model_group_alias( + requested_model=self.data.get("model"), + router_settings=router_settings, + user_api_key_dict=user_api_key_dict, + llm_router=llm_router, + ) + if alias_target is not None: + self.data["model"] = alias_target + self.data["litellm_call_id"] = request.headers.get("x-litellm-call-id", str(uuid.uuid4())) DDSpanTagger.tag_call_id(self.data.get("litellm_call_id")) DDSpanTagger.tag_request( @@ -1339,23 +1403,6 @@ class ProxyBaseLLMRequestProcessing: call_type=route_type, # type: ignore ) - # Apply hierarchical router_settings (Key > Team) - # Global router_settings are already on the Router object itself. - if llm_router is not None and proxy_config is not None: - from litellm.proxy.proxy_server import prisma_client - - router_settings = await proxy_config._get_hierarchical_router_settings( - user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, - proxy_logging_obj=proxy_logging_obj, - ) - - # If router_settings found (from key or team), apply them - # Pass settings as per-request overrides instead of creating a new Router - # This avoids expensive Router instantiation on each request - if router_settings is not None: - self.data["router_settings_override"] = router_settings - if "messages" in self.data and self.data["messages"]: logging_obj.update_messages(self.data["messages"]) diff --git a/litellm/router.py b/litellm/router.py index e6613e1d302..6bf1bdfc670 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -115,6 +115,7 @@ from litellm.router_utils.common_utils import ( _is_proxy_admin_request, filter_team_based_models, filter_web_search_deployments, + resolve_model_group_alias, ) from litellm.router_utils.cooldown_cache import CooldownCache from litellm.router_utils.cooldown_handlers import ( @@ -10331,16 +10332,7 @@ class Router: - str, the litellm model name - None, if model is not in model group alias """ - if model not in self.model_group_alias: - return None - - _item = self.model_group_alias[model] - if isinstance(_item, str): - model = _item - else: - model = _item["model"] - - return model + return resolve_model_group_alias(self.model_group_alias, model) def _get_deployment_by_litellm_model(self, model: str) -> list: """ diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index 189296a8955..bf1c814c049 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -22,6 +22,27 @@ def _is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool return getattr(user_api_key_auth, "user_role", None) == "proxy_admin" +def resolve_model_group_alias(model_group_alias: object, model: str) -> str | None: + """ + Resolve ``model`` through a ``model_group_alias`` map. + + Handles both supported entry shapes, the plain string form + ``{"alias": "target"}`` and the item form + ``{"alias": {"model": "target", "hidden": true}}``, and tolerates malformed + entries: the map can come from a key or team row rather than from validated + config, so a bad value must not raise mid-request. + + Returns the target model group, or None when the map does not rewrite ``model``. + """ + if not isinstance(model_group_alias, Mapping): + return None + entry = model_group_alias.get(model) + target = entry.get("model") if isinstance(entry, Mapping) else entry + if not isinstance(target, str) or not target: + return None + return target + + def get_litellm_params_sensitive_credential_hash(litellm_params: dict) -> str: """ Hash of the credential params, used for mapping the file id to the right model diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 9a3702d2355..28d4d87e26f 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -2016,6 +2016,48 @@ async def test_ProxyConfig__get_hierarchical_router_settings_missing_returns_non assert out is None +@pytest.mark.asyncio +async def test_ProxyConfig__get_hierarchical_router_settings_falls_back_to_team(monkeypatch): + """A key with no router_settings inherits the team's, so a team-level + model_group_alias reaches the request path at all.""" + pc = ProxyConfig() + fake_key = SimpleNamespace(router_settings=None, team_id="team-1") + team_settings = {"model_group_alias": {"group-a": "group-b"}} + monkeypatch.setattr( + "litellm.proxy.proxy_server.get_team_object", + AsyncMock(return_value=SimpleNamespace(router_settings=team_settings)), + ) + + out = await pc._get_hierarchical_router_settings( + user_api_key_dict=fake_key, + prisma_client=None, + proxy_logging_obj=None, + ) + + assert out == team_settings + + +@pytest.mark.asyncio +async def test_ProxyConfig__get_hierarchical_router_settings_key_shadows_team_entirely(monkeypatch): + """Resolution returns whichever object it finds first, it does not merge + per field, so a key that sets any router setting hides every team setting + including an alias the key itself never set.""" + pc = ProxyConfig() + fake_key = SimpleNamespace(router_settings={"num_retries": 3}, team_id="team-1") + team_lookup = AsyncMock(return_value=SimpleNamespace(router_settings={"model_group_alias": {"group-a": "group-b"}})) + monkeypatch.setattr("litellm.proxy.proxy_server.get_team_object", team_lookup) + + out = await pc._get_hierarchical_router_settings( + user_api_key_dict=fake_key, + prisma_client=None, + proxy_logging_obj=None, + ) + + assert out == {"num_retries": 3} + assert "model_group_alias" not in out + team_lookup.assert_not_called() + + # --------------------------------------------------------------------------- # ProxyConfig._add_router_settings_from_db_config # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 3bb84e095a0..4d98a05da8d 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -1,6 +1,7 @@ import asyncio import copy import datetime +from types import SimpleNamespace from typing import AsyncGenerator, Callable, Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -28,11 +29,14 @@ from litellm.proxy.common_request_processing import ( _is_azure_model_router_request, _override_openai_response_model, _parse_event_data_for_error, + _resolve_per_request_model_group_alias, _should_return_raw_model_name, _UpstreamClosingStreamingResponse, create_response, ) from litellm.proxy.dd_span_tagger import DDSpanTagger +from litellm.proxy._types import ProxyException +from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth from litellm.proxy.utils import ProxyLogging @@ -5354,3 +5358,228 @@ class TestModelDeploymentsSupportStreamOptions: def test_non_string_model_is_not_injected(self): assert self._support(None, None) is False + + +class TestPerRequestModelGroupAlias: + """``router_settings.model_group_alias`` on a key or team has to be resolved + by the proxy: the Router resolves aliases from its own shared instance + attribute, which only ever holds the global config map.""" + + @staticmethod + def _router() -> litellm.Router: + return litellm.Router( + model_list=[ + { + "model_name": "group-a", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}, + }, + { + "model_name": "group-b", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-fake"}, + }, + ] + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "alias_map, expected", + [ + ({"group-a": "group-b"}, "group-b"), + ({"group-a": {"model": "group-b", "hidden": True}}, "group-b"), + ({"group-b": "group-a"}, None), + ({"group-a": "group-a"}, None), + ({"group-a": {"hidden": True}}, None), + ({}, None), + (None, None), + ], + ) + async def test_resolves_alias_for_the_requested_model_group(self, alias_map, expected): + resolved = await _resolve_per_request_model_group_alias( + requested_model="group-a", + router_settings={"model_group_alias": alias_map}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", models=[]), + llm_router=self._router(), + ) + + assert resolved == expected + + @pytest.mark.asyncio + async def test_alias_target_outside_the_key_allowlist_is_rejected(self): + """Access was authorized against the requested group, so a rewrite that + the key could not have requested directly must not be served.""" + with pytest.raises(ProxyException) as exc_info: + await _resolve_per_request_model_group_alias( + requested_model="group-a", + router_settings={"model_group_alias": {"group-a": "group-b"}}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", models=["group-a"]), + llm_router=self._router(), + ) + + assert exc_info.value.code == "403" + assert "group-b" in exc_info.value.message + + @pytest.mark.asyncio + async def test_alias_target_inside_the_key_allowlist_resolves(self): + resolved = await _resolve_per_request_model_group_alias( + requested_model="group-a", + router_settings={"model_group_alias": {"group-a": "group-b"}}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", models=["group-a", "group-b"]), + llm_router=self._router(), + ) + + assert resolved == "group-b" + + @pytest.mark.asyncio + @pytest.mark.parametrize("requested_model", [None, ["group-a", "group-b"]]) + async def test_non_string_requested_model_is_left_alone(self, requested_model): + """The routed model is not always a string (a batch request carries a + list), and an unhashable one must not blow up the alias lookup.""" + resolved = await _resolve_per_request_model_group_alias( + requested_model=requested_model, + router_settings={"model_group_alias": {"group-a": "group-b"}}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", models=[]), + llm_router=self._router(), + ) + + assert resolved is None + + @pytest.mark.asyncio + async def test_pre_call_logic_rewrites_the_requested_model(self, monkeypatch): + """End to end through the request path: a key carrying the alias must + leave pre-call processing pointing at the alias target, not at the + group the caller asked for.""" + processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "group-a"}) + mock_request = MagicMock(spec=Request) + mock_request.headers = {} + + async def mock_add_litellm_data_to_request(*args, **kwargs): + return kwargs.get("data", {}) + + async def passthrough_pre_call_hook(user_api_key_dict, data, call_type): + return copy.deepcopy(data) + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=passthrough_pre_call_hook) + monkeypatch.setattr( + litellm.proxy.common_request_processing, + "add_litellm_data_to_request", + mock_add_litellm_data_to_request, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock()) + + mock_proxy_config = MagicMock(spec=ProxyConfig) + mock_proxy_config._get_hierarchical_router_settings = AsyncMock( + return_value={"model_group_alias": {"group-a": "group-b"}} + ) + + returned_data, _ = await processing_obj.common_processing_pre_call_logic( + request=mock_request, + general_settings={}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", models=[]), + proxy_logging_obj=mock_proxy_logging_obj, + proxy_config=mock_proxy_config, + route_type="acompletion", + llm_router=self._router(), + ) + + assert returned_data["model"] == "group-b" + assert returned_data["router_settings_override"] == {"model_group_alias": {"group-a": "group-b"}} + # The rewrite has to land before the pre-call hooks: they are where + # per-model budgets and rate limits are enforced, so resolving later + # applies the requested group's limits to a call the target serves. + assert mock_proxy_logging_obj.pre_call_hook.call_args.kwargs["data"]["model"] == "group-b" + + @pytest.mark.asyncio + async def test_team_level_alias_rewrites_the_requested_model(self, monkeypatch): + """The team path is separate resolution, not a variant of the key path: + settings are looked up on the team only when the key carries none. Runs + the real hierarchical lookup rather than mocking it, so this covers the + team half of the fix end to end.""" + from litellm.proxy.proxy_server import ProxyConfig as RealProxyConfig + + processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "group-a"}) + mock_request = MagicMock(spec=Request) + mock_request.headers = {} + + async def mock_add_litellm_data_to_request(*args, **kwargs): + return kwargs.get("data", {}) + + async def passthrough_pre_call_hook(user_api_key_dict, data, call_type): + return copy.deepcopy(data) + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=passthrough_pre_call_hook) + monkeypatch.setattr( + litellm.proxy.common_request_processing, + "add_litellm_data_to_request", + mock_add_litellm_data_to_request, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock()) + monkeypatch.setattr( + "litellm.proxy.proxy_server.get_team_object", + AsyncMock(return_value=SimpleNamespace(router_settings={"model_group_alias": {"group-a": "group-b"}})), + ) + + returned_data, _ = await processing_obj.common_processing_pre_call_logic( + request=mock_request, + general_settings={}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", models=[], team_id="team-1"), + proxy_logging_obj=mock_proxy_logging_obj, + proxy_config=RealProxyConfig(), + route_type="acompletion", + llm_router=self._router(), + ) + + assert returned_data["model"] == "group-b" + + @pytest.mark.asyncio + async def test_model_level_guardrails_resolve_against_the_alias_target(self, monkeypatch): + """Model-level guardrails are merged by model group name, so the merge + must see the target rather than the group the caller named.""" + processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "group-a"}) + mock_request = MagicMock(spec=Request) + mock_request.headers = {} + + async def mock_add_litellm_data_to_request(*args, **kwargs): + return kwargs.get("data", {}) + + async def passthrough_pre_call_hook(user_api_key_dict, data, call_type): + return copy.deepcopy(data) + + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=passthrough_pre_call_hook) + monkeypatch.setattr( + litellm.proxy.common_request_processing, + "add_litellm_data_to_request", + mock_add_litellm_data_to_request, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock()) + + merged_for: list = [] + + def recording_merge(data, llm_router, trust_client_model_info=True): + merged_for.append(data.get("model")) + return data + + monkeypatch.setattr( + litellm.proxy.common_request_processing, + "_check_and_merge_model_level_guardrails", + recording_merge, + ) + + mock_proxy_config = MagicMock(spec=ProxyConfig) + mock_proxy_config._get_hierarchical_router_settings = AsyncMock( + return_value={"model_group_alias": {"group-a": "group-b"}} + ) + + await processing_obj.common_processing_pre_call_logic( + request=mock_request, + general_settings={}, + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="hash", models=[]), + proxy_logging_obj=mock_proxy_logging_obj, + proxy_config=mock_proxy_config, + route_type="acompletion", + llm_router=self._router(), + ) + + assert merged_for == ["group-b"] diff --git a/tests/test_litellm/proxy/test_model_level_guardrails.py b/tests/test_litellm/proxy/test_model_level_guardrails.py index 48163bf5ed5..a1278e399b5 100644 --- a/tests/test_litellm/proxy/test_model_level_guardrails.py +++ b/tests/test_litellm/proxy/test_model_level_guardrails.py @@ -598,10 +598,15 @@ async def test_pre_call_merges_model_level_guardrails_before_pre_call_hook(): } ) - captured_pre_call_data: dict = {} + captured_pre_call_guardrails: list = [] async def fake_pre_call_hook(*, user_api_key_dict, data, call_type): - captured_pre_call_data.update(data) + # Snapshot the list rather than the dict: metadata is shared by + # reference, so a merge that happens after this point would otherwise + # show up here retroactively and the assertion would pass either way. + captured_pre_call_guardrails.extend( + (data.get("metadata") or {}).get("guardrails") or data.get("guardrails") or [] + ) return data proxy_logging = MagicMock() @@ -616,13 +621,9 @@ async def test_pre_call_merges_model_level_guardrails_before_pre_call_hook(): proxy_config = MagicMock() proxy_config._get_hierarchical_router_settings = AsyncMock(return_value=None) - # Stop the function before any post-pre_call_hook logic so we can keep - # the test focused. Raising _StopAfterPreCall in the next await fires - # right after the guardrail merge + pre_call_hook complete. - class _StopAfterPreCall(Exception): - pass - - proxy_config._get_hierarchical_router_settings.side_effect = _StopAfterPreCall() + # Assert on what pre_call_hook was handed rather than short-circuiting the + # function part way through: a sentinel keyed to one particular later call + # silently stops testing the ordering as soon as that call moves. with ( patch( @@ -640,30 +641,24 @@ async def test_pre_call_merges_model_level_guardrails_before_pre_call_hook(): ): from litellm.proxy._types import UserAPIKeyAuth - try: - await processing.common_processing_pre_call_logic( - request=MagicMock(headers={}, url=MagicMock(path="/v1/chat/completions")), - general_settings={}, - user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), - proxy_logging_obj=proxy_logging, - proxy_config=proxy_config, - route_type="acompletion", - version=None, - user_model=None, - user_temperature=None, - user_request_timeout=None, - user_max_tokens=None, - user_api_base=None, - model=None, - llm_router=mock_router, - ) - except _StopAfterPreCall: - pass + await processing.common_processing_pre_call_logic( + request=MagicMock(headers={}, url=MagicMock(path="/v1/chat/completions")), + general_settings={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test-key"), + proxy_logging_obj=proxy_logging, + proxy_config=proxy_config, + route_type="acompletion", + version=None, + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + model=None, + llm_router=mock_router, + ) # The pre_call_hook must have received data with the model-level # guardrail already merged in. Before the fix, this assertion fails # because pre_call_hook saw the original data without merge. - merged = (captured_pre_call_data.get("metadata") or {}).get("guardrails") or ( - captured_pre_call_data.get("guardrails") or [] - ) - assert "my-pre-call-guardrail" in merged + assert "my-pre-call-guardrail" in captured_pre_call_guardrails diff --git a/tests/test_litellm/router_utils/test_router_utils_common_utils.py b/tests/test_litellm/router_utils/test_router_utils_common_utils.py index efa5f2382dc..7d453c72652 100644 --- a/tests/test_litellm/router_utils/test_router_utils_common_utils.py +++ b/tests/test_litellm/router_utils/test_router_utils_common_utils.py @@ -10,6 +10,7 @@ from litellm.router_utils.common_utils import ( add_model_file_id_mappings, filter_team_based_models, filter_web_search_deployments, + resolve_model_group_alias, ) @@ -516,3 +517,44 @@ class TestAddModelFileIdMappings: def test_should_return_empty_mapping_when_given_empty_list(self): result = add_model_file_id_mappings([], []) assert result == {} + + +class TestResolveModelGroupAlias: + """``model_group_alias`` maps reach this helper from validated config and + from key/team rows, so both entry shapes must resolve and malformed entries + must not raise mid-request.""" + + @pytest.mark.parametrize( + "alias_map, expected", + [ + ({"group-a": "group-b"}, "group-b"), + ({"group-a": {"model": "group-b", "hidden": True}}, "group-b"), + ({"group-a": {"model": "group-b"}}, "group-b"), + ({"other": "group-b"}, None), + ({}, None), + (None, None), + ("not-a-map", None), + ({"group-a": {"hidden": True}}, None), + ({"group-a": {"model": 5}}, None), + ({"group-a": 5}, None), + ({"group-a": None}, None), + ({"group-a": ""}, None), + ], + ) + def test_resolves_both_entry_shapes_and_tolerates_malformed_entries(self, alias_map, expected): + assert resolve_model_group_alias(alias_map, "group-a") == expected + + def test_router_alias_resolution_uses_the_shared_helper(self): + router = Router( + model_list=[ + { + "model_name": "group-b", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake"}, + } + ], + model_group_alias={"group-a": "group-b", "group-item": {"model": "group-b", "hidden": True}}, + ) + + assert router._get_model_from_alias("group-a") == "group-b" + assert router._get_model_from_alias("group-item") == "group-b" + assert router._get_model_from_alias("group-b") is None