diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 6922b5f6cbd..2944587c3e1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -5366,20 +5366,26 @@ def can_customer_access_model( valid_token: UserAPIKeyAuth | None, ) -> Literal[True]: team_model_aliases: Final = team_model_aliases_for_auth_check(valid_token) if valid_token is not None else None - listed_names: Final = frozenset(end_user_object.models or ()) - unlisted_aliases: Final = ( - MappingProxyType({alias: target for alias, target in team_model_aliases.items() if alias not in listed_names}) - if team_model_aliases - else None - ) team_id: Final = valid_token.team_id if valid_token is not None else None - return _can_object_call_model( - model=_resolve_team_alias(model, unlisted_aliases, team_id, llm_router), - llm_router=llm_router, - models=end_user_object.models, - key_model_aliases=key_model_aliases_for_auth_check(valid_token), - object_type="customer", - ) + key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token) + + def check(name: str) -> None: + team_target: Final = ( + _live_team_alias_target(name, team_model_aliases, team_id, llm_router) if team_model_aliases else name + ) + if team_target != name and name in (end_user_object.models or ()): + return + _can_object_call_model( + model=team_target, + llm_router=llm_router, + models=end_user_object.models, + key_model_aliases=key_model_aliases, + object_type="customer", + ) + + for name in (model,) if isinstance(model, str) else model: + check(name) + return True async def can_user_call_model( diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index 6f62e6ed092..63832c485f9 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -8882,6 +8882,9 @@ async def _common_checks_for_customer_model( customer_models: list[str], request_overrides: Mapping[str, object] | None = None, team_model_aliases: dict[str, str] | None = None, + key_model_aliases: dict[str, str] | None = None, + team_id: str | None = None, + llm_router: "Router | None" = None, ) -> bool: from litellm.proxy.auth.auth_checks import common_checks @@ -8897,9 +8900,14 @@ async def _common_checks_for_customer_model( global_proxy_spend=None, general_settings={}, route="/chat/completions", - llm_router=None, + llm_router=llm_router, proxy_logging_obj=MagicMock(), - valid_token=UserAPIKeyAuth(token="test-token", team_model_aliases=team_model_aliases), + valid_token=UserAPIKeyAuth( + token="test-token", + team_id=team_id, + team_model_aliases=team_model_aliases, + aliases=key_model_aliases or {}, + ), request=MagicMock(spec=Request), skip_budget_checks=True, ) @@ -8958,6 +8966,69 @@ async def test_common_checks_matches_team_alias_target_against_customer_allowlis ) +@pytest.mark.asyncio +async def test_common_checks_prefers_team_alias_over_same_named_key_alias_for_customer() -> None: + team_model_aliases: Final = {"fast": "m1"} + key_model_aliases: Final = {"fast": "m2"} + + for customer_models in (["fast"], ["m1"]): + assert ( + await _common_checks_for_customer_model( + model="fast", + customer_models=customer_models, + team_model_aliases=team_model_aliases, + key_model_aliases=key_model_aliases, + ) + is True + ) + with pytest.raises(ModelAccessDeniedProxyException) as exc_info: + await _common_checks_for_customer_model( + model="fast", + customer_models=["m2"], + team_model_aliases=team_model_aliases, + key_model_aliases=key_model_aliases, + ) + + assert exc_info.value.type == ProxyErrorTypes.customer_model_access_denied + + +@pytest.mark.asyncio +async def test_common_checks_applies_key_alias_for_customer_when_team_alias_target_is_deleted() -> None: + from litellm import Router + + llm_router: Final = Router( + model_list=[ + {"model_name": name, "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-api-key"}} + for name in ("m1", "m2") + ] + ) + team_model_aliases: Final = {"fast": "model_name_team-1_deleted"} + key_model_aliases: Final = {"fast": "m2"} + + assert ( + await _common_checks_for_customer_model( + model="fast", + customer_models=["m2"], + team_model_aliases=team_model_aliases, + key_model_aliases=key_model_aliases, + team_id="team-1", + llm_router=llm_router, + ) + is True + ) + with pytest.raises(ModelAccessDeniedProxyException) as exc_info: + await _common_checks_for_customer_model( + model="fast", + customer_models=["fast"], + team_model_aliases=team_model_aliases, + key_model_aliases=key_model_aliases, + team_id="team-1", + llm_router=llm_router, + ) + + assert exc_info.value.type == ProxyErrorTypes.customer_model_access_denied + + @pytest.mark.parametrize( ("model", "customer_models", "denied"), (