mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
a1a42768c1
commit
d9b1d57f07
2 changed files with 92 additions and 15 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue