mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix
This commit is contained in:
parent
1c889f0e44
commit
aa311bc833
2 changed files with 108 additions and 5 deletions
|
|
@ -2873,6 +2873,40 @@ def _model_in_team_aliases(
|
|||
return False
|
||||
|
||||
|
||||
def _model_list_allows_model(model: str, allowed_models: List[str]) -> bool:
|
||||
if (
|
||||
"*" in allowed_models
|
||||
or SpecialModelNames.all_proxy_models.value in allowed_models
|
||||
):
|
||||
return True
|
||||
if model in allowed_models:
|
||||
return True
|
||||
return _model_matches_any_wildcard_pattern_in_list(
|
||||
model=model, allowed_model_list=allowed_models
|
||||
)
|
||||
|
||||
|
||||
def _filter_team_model_aliases_by_allowed_models(
|
||||
team_model_aliases: Optional[Dict[str, str]],
|
||||
allowed_models: List[str],
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""
|
||||
Return only aliases whose resolved target is allowed by `allowed_models`.
|
||||
|
||||
This is used when the effective allowlist comes from access groups. Without
|
||||
this filter, any team alias name would bypass the access-group model list.
|
||||
"""
|
||||
if not team_model_aliases:
|
||||
return team_model_aliases
|
||||
|
||||
filtered_aliases = {
|
||||
alias: target
|
||||
for alias, target in team_model_aliases.items()
|
||||
if _model_list_allows_model(model=target, allowed_models=allowed_models)
|
||||
}
|
||||
return filtered_aliases
|
||||
|
||||
|
||||
async def can_key_call_model(
|
||||
model: Union[str, List[str]],
|
||||
llm_model_list: Optional[list],
|
||||
|
|
@ -2906,7 +2940,10 @@ async def can_key_call_model(
|
|||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=models_from_groups,
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
team_model_aliases=_filter_team_model_aliases_by_allowed_models(
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
allowed_models=models_from_groups,
|
||||
),
|
||||
team_id=valid_token.team_id,
|
||||
object_type="key",
|
||||
)
|
||||
|
|
@ -2945,7 +2982,10 @@ async def can_key_call_model(
|
|||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=models_from_groups,
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
team_model_aliases=_filter_team_model_aliases_by_allowed_models(
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
allowed_models=models_from_groups,
|
||||
),
|
||||
team_id=valid_token.team_id,
|
||||
object_type="key",
|
||||
)
|
||||
|
|
@ -3006,7 +3046,10 @@ async def can_team_access_model(
|
|||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=models_from_groups,
|
||||
team_model_aliases=team_model_aliases,
|
||||
team_model_aliases=_filter_team_model_aliases_by_allowed_models(
|
||||
team_model_aliases=team_model_aliases,
|
||||
allowed_models=models_from_groups,
|
||||
),
|
||||
team_id=team_object.team_id if team_object else None,
|
||||
object_type="team",
|
||||
)
|
||||
|
|
@ -3070,12 +3113,16 @@ async def _key_access_group_grants_model(
|
|||
|
||||
if not authorized_models:
|
||||
return False
|
||||
authorized_models = list(set(authorized_models))
|
||||
try:
|
||||
_can_object_call_model(
|
||||
model=model,
|
||||
llm_router=llm_router,
|
||||
models=list(set(authorized_models)),
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
models=authorized_models,
|
||||
team_model_aliases=_filter_team_model_aliases_by_allowed_models(
|
||||
team_model_aliases=valid_token.team_model_aliases,
|
||||
allowed_models=authorized_models,
|
||||
),
|
||||
team_id=valid_token.team_id,
|
||||
object_type="key",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1269,6 +1269,62 @@ async def test_can_key_call_model_denies_when_access_group_ids_resolve_no_models
|
|||
assert "resolved to no model permissions" in warning.call_args.args[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_can_key_call_model_denies_team_alias_outside_access_group_models():
|
||||
"""Team aliases do not bypass a key's access-group model allowlist."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_model
|
||||
|
||||
user_api_key_object = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
models=[],
|
||||
access_group_ids=["ag-with-gpt4"],
|
||||
team_model_aliases={"premium-alias": "claude-3"},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=["gpt-4"],
|
||||
):
|
||||
with pytest.raises(ProxyException):
|
||||
await can_key_call_model(
|
||||
model="premium-alias",
|
||||
llm_model_list=[],
|
||||
valid_token=user_api_key_object,
|
||||
llm_router=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_can_key_call_model_allows_team_alias_inside_access_group_models():
|
||||
"""Team aliases remain usable when their target is in the key's access groups."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy.auth.auth_checks import can_key_call_model
|
||||
|
||||
user_api_key_object = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
models=[],
|
||||
access_group_ids=["ag-with-gpt4"],
|
||||
team_model_aliases={"premium-alias": "gpt-4"},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
|
||||
new_callable=AsyncMock,
|
||||
return_value=["gpt-4"],
|
||||
):
|
||||
await can_key_call_model(
|
||||
model="premium-alias",
|
||||
llm_model_list=[],
|
||||
valid_token=user_api_key_object,
|
||||
llm_router=None,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _key_access_group_grants_model (key access group overriding team restriction)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue