fix(proxy): let a listed team alias win over a same-named key alias in the customer model check (#44677)

* fix(proxy): let a listed team alias win over a same-named key alias in the customer model check

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(proxy): check each requested name in a plain loop in can_customer_access_model

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): only let a listed team alias skip the customer check when its target is live

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(proxy): move the per-name customer alias check into a local function

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: kerry <kerry@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 22:43:16 +00:00 • committed by GitHub
parent a1a42768c1
commit d9b1d57f07
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 92 additions and 15 deletions

View file

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

View file

@ -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"),
(