diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 67950e603c0..f19a8055ae6 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -89,6 +89,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, ) +from litellm.proxy.common_utils.model_listing_utils import alias_map from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import ( END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL, @@ -722,6 +723,7 @@ async def _run_project_checks( model=_model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if not skip_budget_checks: @@ -1018,6 +1020,7 @@ async def common_checks( team_object=team_object, llm_router=llm_router, team_model_aliases=(valid_token.team_model_aliases if valid_token else None), + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -1027,6 +1030,7 @@ async def common_checks( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -1043,6 +1047,7 @@ async def common_checks( proxy_logging_obj=proxy_logging_obj, team_membership=loaded_team_membership, team_membership_loaded=team_membership_loaded, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent @@ -1081,6 +1086,7 @@ async def common_checks( model=_model, llm_router=llm_router, user_object=user_object, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget) @@ -4349,6 +4355,7 @@ def _can_object_call_model( models: list[str], team_model_aliases: dict[str, str] | None = None, team_id: str | None = None, + key_model_aliases: Mapping[str, str] | None = None, object_type: Literal["user", "team", "key", "org", "project", "agent"] = "user", fallback_depth: int = 0, ) -> Literal[True]: @@ -4378,6 +4385,7 @@ def _can_object_call_model( models=models, team_model_aliases=team_model_aliases, team_id=team_id, + key_model_aliases=key_model_aliases, object_type=object_type, fallback_depth=fallback_depth + 1, ) @@ -4386,13 +4394,32 @@ def _can_object_call_model( from litellm.router_strategy.complexity_router.context_compaction import native_compaction_parent compaction_parent: Final = native_compaction_parent(model) - potential_models: Final = [model, compaction_parent] if compaction_parent is not None else [model] - if model in litellm.model_alias_map: - potential_models.append(litellm.model_alias_map[model]) - elif llm_router and model in llm_router.model_group_alias: - _model: Final = llm_router._get_model_from_alias(model) - if _model: - potential_models.append(_model) + global_or_router_alias_target: Final = ( + litellm.model_alias_map[model] + if model in litellm.model_alias_map + else ( + llm_router._get_model_from_alias(model) + if llm_router is not None and model in llm_router.model_group_alias + else None + ) + ) + after_team_alias: Final = team_model_aliases.get(model, model) if team_model_aliases else model + after_key_alias: Final = ( + key_model_aliases.get(after_team_alias, after_team_alias) if key_model_aliases else after_team_alias + ) + after_global_alias: Final = litellm.model_alias_map.get(after_key_alias, after_key_alias) + dispatched_model: Final = ( + key_model_aliases.get(after_global_alias, after_global_alias) if key_model_aliases else after_global_alias + ) + key_alias_applied: Final = after_key_alias != after_team_alias or dispatched_model != after_global_alias + potential_models: Final = ( + (dispatched_model,) + if key_alias_applied + else ( + *((model, compaction_parent) if compaction_parent is not None else (model,)), + *((global_or_router_alias_target,) if global_or_router_alias_target else ()), + ) + ) ## check model access for alias + underlying model - allow if either is in allowed models for m in potential_models: @@ -4418,6 +4445,35 @@ def _can_object_call_model( ) +def _resolve_team_alias( + model: str | list[str], + team_model_aliases: dict[str, str] | None, + team_id: str | None, + llm_router: Router | None, +) -> str | list[str]: + if not team_model_aliases: + return model + if isinstance(model, str): + return _live_team_alias_target(model, team_model_aliases, team_id, llm_router) + return [ # mutable-ok: _can_object_call_model takes list[str] + _live_team_alias_target(name, team_model_aliases, team_id, llm_router) for name in model + ] + + +def _live_team_alias_target( + model: str, team_model_aliases: dict[str, str], team_id: str | None, llm_router: Router | None +) -> str: + target: Final = team_model_aliases.get(model) + if target is None: + return model + deleted_team_deployment: Final = ( + llm_router is not None + and target.startswith(f"model_name_{team_id}_") + and target not in llm_router.model_name_to_deployment_indices + ) + return model if deleted_team_deployment else target + + async def _check_agent_access_group_model_access( model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str] valid_token: UserAPIKeyAuth | None, @@ -4438,12 +4494,14 @@ async def _check_agent_access_group_model_access( param="model", code=status.HTTP_403_FORBIDDEN, ) + dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router) return _can_object_call_model( - model=model, + model=dispatched, llm_router=llm_router, models=sorted(ceiling.models), team_id=valid_token.team_id, object_type="agent", + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4471,12 +4529,14 @@ async def _check_agent_caller_model_access( if caller_auth is None: return caller_team: Final = await load_team(valid_token) + caller_key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token) if caller_team is not None: await can_team_access_model( model=model, team_object=caller_team, llm_router=llm_router, prisma_client=prisma_client, + key_model_aliases=caller_key_model_aliases, ) await _check_team_member_model_access( model=model, @@ -4486,12 +4546,18 @@ async def _check_agent_caller_model_access( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=caller_key_model_aliases, ) return caller_user: Final = await load_user(valid_token) if caller_user is None: return - await can_user_call_model(model=model, llm_router=llm_router, user_object=caller_user) + await can_user_call_model( + model=model, + llm_router=llm_router, + user_object=caller_user, + key_model_aliases=caller_key_model_aliases, + ) def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None = None) -> bool: @@ -4512,6 +4578,10 @@ def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None return False +def key_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth | None) -> Mapping[str, str] | None: + return alias_map(valid_token.aliases) if valid_token is not None and valid_token.aliases else None + + def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]: """ Expand key model sentinels before auth checks. @@ -4831,6 +4901,7 @@ async def can_key_call_model( models=key_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) except ProxyException: @@ -4848,6 +4919,7 @@ async def can_key_call_model( models=models_from_groups, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) raise @@ -4906,6 +4978,7 @@ async def can_key_call_resolved_model( team_object=team_object, llm_router=llm_router, team_model_aliases=valid_token.team_model_aliases, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -4915,6 +4988,7 @@ async def can_key_call_resolved_model( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -4927,6 +5001,7 @@ async def can_key_call_resolved_model( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if valid_token.project_id is not None: @@ -4941,6 +5016,7 @@ async def can_key_call_resolved_model( model=model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4968,6 +5044,7 @@ async def can_team_access_model( team_object: LiteLLM_TeamTable | None, llm_router: Router | None, team_model_aliases: dict[str, str] | None = None, + key_model_aliases: Mapping[str, str] | None = None, prisma_client: DatabaseClient | None = None, ) -> Literal[True]: """ @@ -4983,6 +5060,7 @@ async def can_team_access_model( models=team_object.models if team_object else [], team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) except ProxyException: @@ -5000,6 +5078,7 @@ async def can_team_access_model( models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])), team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) raise @@ -5058,6 +5137,7 @@ async def _key_access_group_grants_model( valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> bool: """ Returns True if the key's `access_group_ids` expand to models that grant @@ -5078,6 +5158,7 @@ async def _key_access_group_grants_model( models=authorized_models, team_model_aliases=valid_token.team_model_aliases if valid_token else None, team_id=valid_token.team_id if valid_token else None, + key_model_aliases=key_model_aliases, object_type="key", ) return True @@ -5089,6 +5170,7 @@ def can_project_access_model( model: str | list[str], project_object: LiteLLM_ProjectTable, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: """ Returns True if the project can access a specific model. @@ -5099,6 +5181,7 @@ def can_project_access_model( model=model, llm_router=llm_router, models=project_object.models if project_object else [], + key_model_aliases=key_model_aliases, object_type="project", ) @@ -5107,6 +5190,7 @@ async def can_user_call_model( model: str | list[str], llm_router: Router | None, user_object: LiteLLM_UserTable | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: if user_object is None: return True @@ -5128,6 +5212,7 @@ async def can_user_call_model( model=model, llm_router=llm_router, models=user_object.models, + key_model_aliases=key_model_aliases, object_type="user", ) @@ -5682,6 +5767,7 @@ async def _check_team_member_model_access( proxy_logging_obj: ProxyLogging, team_membership: LiteLLM_TeamMembership | None = None, team_membership_loaded: bool = False, + key_model_aliases: Mapping[str, str] | None = None, ) -> None: """ Check if a team member's per-member model scope allows access to the requested model. @@ -5717,6 +5803,7 @@ async def _check_team_member_model_access( models=member_allowed_models, object_type="team", team_id=team_object.team_id, + key_model_aliases=key_model_aliases, ) except ProxyException: internal_message: Final = ( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ae95e94dd2d..22c3a248b9d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_checks import ( get_user_object, is_valid_fallback_model, jwt_key_mapping_cache_key, + key_model_aliases_for_auth_check, resolve_and_validate_end_user_id, resolve_default_end_user_budget, ) @@ -469,6 +470,7 @@ async def _check_key_model_budget_with_fallback( models=valid_token.team_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="team", ) except ProxyException: diff --git a/litellm/proxy/common_utils/model_listing_utils.py b/litellm/proxy/common_utils/model_listing_utils.py index 3c6555662e6..8958fb20918 100644 --- a/litellm/proxy/common_utils/model_listing_utils.py +++ b/litellm/proxy/common_utils/model_listing_utils.py @@ -180,7 +180,7 @@ def caller_alias_maps( return CallerAliases((team_aliases, key_aliases), (team_aliases, key_aliases, litellm.model_alias_map, key_aliases)) -def _alias_map(aliases: object) -> Mapping[str, str]: +def alias_map(aliases: object) -> Mapping[str, str]: try: entries: Final = _ALIAS_ENTRIES.validate_python(aliases, strict=True) except ValidationError: @@ -204,7 +204,7 @@ def alias_target(model_id: str, aliases: CallerAliases, listed: Container[str] = already `listed` keeps its own row, so it is never rewritten.""" if model_id in listed: return None - return _rewrite(model_id, tuple(_alias_map(alias_map) for alias_map in aliases.rewrite)) + return _rewrite(model_id, tuple(alias_map(raw) for raw in aliases.rewrite)) def alias_listing_entries( @@ -213,8 +213,8 @@ def alias_listing_entries( ) -> tuple[tuple[str, str], ...]: """`entries` plus one `(alias, lookup_id)` row per key or team alias whose target is listed. An alias colliding with a listed id keeps the listed entry.""" - maps: Final = tuple(_alias_map(alias_map) for alias_map in aliases.rewrite) - own: Final = tuple(_alias_map(alias_map) for alias_map in aliases.own) + maps: Final = tuple(alias_map(raw) for raw in aliases.rewrite) + own: Final = tuple(alias_map(raw) for raw in aliases.own) lookup_by_response: Final = MappingProxyType(dict(entries)) lookup_ids: Final = frozenset(lookup_by_response.values()) targets: Final = MappingProxyType( diff --git a/tests/integration/authorization/test_key_alias_model_access.py b/tests/integration/authorization/test_key_alias_model_access.py new file mode 100644 index 00000000000..50fc53bd4a9 --- /dev/null +++ b/tests/integration/authorization/test_key_alias_model_access.py @@ -0,0 +1,72 @@ +import uuid +from typing import Final + +import httpx + +from tests.integration._support.client import Gateway, eventually, object_value, string_value + + +def _listed_model_ids(response: httpx.Response) -> frozenset[str]: + entries: Final = response.json()["data"] + assert isinstance(entries, list), response.text + return frozenset(string_value(object_value(entry)["id"]) for entry in entries) + + +def _listed_and_callable(gateway: Gateway, key: str, model: str, alias: str) -> None: + """Every id /v1/models lists for this key must be callable by the same key.""" + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({model, alias}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + listed: Final = _listed_model_ids(response) + assert listed == frozenset({model, alias}), response.text + for model_id in sorted(listed): + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model_id, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 200, f"listed id {model_id} is not callable: {called.status_code} {called.text}" + + +def test_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[model], aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_team_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + team_id: Final = scenario.team(models=[model]) + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(team_id=team_id, aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_key_alias_to_model_outside_key_allowlist_is_hidden_and_denied(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + allowed: Final = scenario.model() + hidden: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[allowed], aliases={alias: hidden}) + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({allowed}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + assert _listed_model_ids(response) == frozenset({allowed}), response.text + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": alias, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 403, called.text + assert "key_model_access_denied" in called.text, called.text diff --git a/tests/test_keys.py b/tests/test_keys.py index 7a5b2502cfd..c1785b88822 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -834,12 +834,12 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint): assert len(model_list["data"]) > 0 if model_access == "gpt-3.5-turbo": if model_endpoint == "/v1/models": - assert ( - len(model_list["data"]) == 1 - ), "model_access={}, model_access_level={}".format( + assert {entry["id"] for entry in model_list["data"]} == { + model_access, + "mistral-7b", + }, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format( model_access, model_access_level ) - assert model_list["data"][0]["id"] == model_access elif model_endpoint == "/model/info": assert isinstance(model_list["data"], list) assert len(model_list["data"]) == 1 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 30f5abdbb98..b811d4453ca 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1783,6 +1783,336 @@ def test_can_object_call_model_access_via_alias_only(): assert result is True +def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): + """A key alias whose target is on the key allowlist resolves like a team alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + result = _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + object_type="key", + fallback_depth=0, + ) + + assert result is True + + +def test_can_object_call_model_key_alias_to_disallowed_target_is_denied(): + """A key alias whose target is outside the key allowlist stays denied.""" + from litellm.proxy._types import ProxyErrorTypes, ProxyException + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + assert exc_info.value.code == "403" + + +@pytest.mark.asyncio +async def test_can_team_access_model_honors_key_alias(): + """A key on a team can call a model through its own alias when the target is on the team allowlist.""" + from litellm.proxy.auth.auth_checks import can_team_access_model + + team_object = LiteLLM_TeamTable( + team_id="team-123", + models=["gpt-4o-mini"], + ) + + assert ( + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + +@pytest.mark.asyncio +async def test_can_key_call_model_honors_key_alias(): + """The real key entry point resolves a key alias to its target before the allowlist check.""" + from litellm.proxy.auth.auth_checks import can_key_call_model + + allowed_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + assert ( + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=allowed_token, + llm_router=None, + ) + is True + ) + + denied_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=denied_token, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch): + """The key alias rewrite precedes the global one at dispatch, so the key target is authorized.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch): + """A key alias on the globally rewritten name resolves the same way the request chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): + """When a key alias fires on the globally rewritten name, only the final target is dispatched.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_name_alone_is_not_enough(): + """A key that may call the alias name but not its target cannot call the alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="bar", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="bar", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_team_alias_applies_before_key_alias(): + """A key alias on the raw name loses to the team alias that rewrites it first at dispatch.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_on_team_alias_target(): + """A key alias on the team-rewritten name resolves like the dispatch chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +@pytest.mark.asyncio +async def test_can_user_call_model_honors_key_alias(): + """A personal-scope key alias resolves to its target before the user allowlist check.""" + from litellm.proxy.auth.auth_checks import can_user_call_model + + user_object = LiteLLM_UserTable(user_id="test-user", models=["gpt-4o-mini"]) + + assert ( + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + ) + + assert exc_info.value.type == ProxyErrorTypes.user_model_access_denied + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_honors_key_alias(): + """A key alias resolves against the member allowlist, not just the raw alias name.""" + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + + membership = LiteLLM_TeamMembership( + user_id="alice", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]), + ) + + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + def test_can_object_call_model_access_via_underlying_model_only(): """ Test that a key can access a model via underlying model even when using an alias. @@ -9139,6 +9469,50 @@ async def test_agent_access_groups_cap_models_even_when_key_allows_them(): assert asked == ["agent-1", "agent-1"] +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_admits_the_key_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5"]) + agent_key.aliases = {"fast": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("fast", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_checks_the_team_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_denies_a_team_alias_outside_the_ceiling(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "claude-sonnet-4-5"} + resolve, _ = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + with pytest.raises(ModelAccessDeniedProxyException) as exc: + await _check_agent_access_group_model_access("foo", agent_key, None, resolve) + assert exc.value.type == ProxyErrorTypes.agent_model_access_denied + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_keeps_the_name_for_a_deleted_team_deployment(): + from litellm.router import Router + + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["foo"]) + agent_key.team_model_aliases = {"foo": "model_name_team-1_deadbeef"} + router: Final = Router(model_list=[]) + resolve, asked = _agent_model_ceiling_resolver(frozenset({"foo"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, router, resolve) is True + assert asked == ["agent-1"] + + @pytest.mark.asyncio async def test_agent_access_groups_naming_no_model_deny_every_model(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[]) @@ -9437,6 +9811,18 @@ async def test_agent_key_acting_for_a_teamless_user_is_capped_at_that_users_mode assert asked == ["team:None", "user:alice", "team:None", "user:alice"] +@pytest.mark.asyncio +async def test_agent_key_alias_resolves_against_the_echoed_teams_models(): + agent_key: Final = _agent_key_acting_for(user_id="alice", team_id="team-a") + agent_key.aliases = {"foo": "bar"} + load_team, load_user, asked = _caller_loaders(LiteLLM_TeamTable(team_id="team-a", models=["bar"]), None) + cache: Final = await _cache_with_membership("alice", "team-a", allowed_models=None) + + await _check_caller_models(agent_key, "foo", load_team, load_user, cache) + + assert asked == ["team:team-a"] + + @pytest.mark.asyncio async def test_agent_key_without_an_echoed_caller_keeps_its_own_models(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"])