mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: align all-team-models sentinel access
This commit is contained in:
parent
5663bcf0d2
commit
d093dc2e3b
4 changed files with 90 additions and 4 deletions
|
|
@ -2949,6 +2949,26 @@ async def _get_agent_ids_from_access_groups(
|
|||
)
|
||||
|
||||
|
||||
def _resolve_all_team_model_sentinel_for_auth_check(
|
||||
models: List[str],
|
||||
llm_router: Optional[Router],
|
||||
team_id: Optional[str],
|
||||
) -> List[str]:
|
||||
if (
|
||||
SpecialModelNames.all_team_models.value not in models
|
||||
or team_id is None
|
||||
or llm_router is None
|
||||
):
|
||||
return models
|
||||
proxy_models = llm_router.get_model_names()
|
||||
non_sentinel_models = [
|
||||
model for model in models if model != SpecialModelNames.all_team_models.value
|
||||
]
|
||||
if not proxy_models:
|
||||
return non_sentinel_models or models
|
||||
return list(dict.fromkeys(non_sentinel_models + proxy_models))
|
||||
|
||||
|
||||
def _check_model_access_helper(
|
||||
model: str,
|
||||
llm_router: Optional[Router],
|
||||
|
|
@ -2966,6 +2986,12 @@ def _check_model_access_helper(
|
|||
model_name=model, team_id=team_id
|
||||
)
|
||||
|
||||
models = _resolve_all_team_model_sentinel_for_auth_check(
|
||||
models=models,
|
||||
llm_router=llm_router,
|
||||
team_id=team_id,
|
||||
)
|
||||
|
||||
if (
|
||||
len(access_groups) > 0 and llm_router is not None
|
||||
): # check if token contains any model access groups
|
||||
|
|
|
|||
|
|
@ -123,11 +123,13 @@ def get_key_models(
|
|||
and user_api_key_dict.team_id is not None
|
||||
):
|
||||
all_models = list(user_api_key_dict.team_models)
|
||||
# GH#30619: if team_models also contains all-team-models,
|
||||
# expand to actual proxy model list instead of leaking
|
||||
# the sentinel string into /model/info
|
||||
if SpecialModelNames.all_team_models.value in all_models:
|
||||
all_models = list(proxy_model_list)
|
||||
all_models = [
|
||||
model
|
||||
for model in all_models
|
||||
if model != SpecialModelNames.all_team_models.value
|
||||
]
|
||||
all_models.extend(proxy_model_list)
|
||||
if include_model_access_groups:
|
||||
all_models.extend(model_access_groups.keys())
|
||||
if SpecialModelNames.all_proxy_models.value in all_models:
|
||||
|
|
|
|||
|
|
@ -351,6 +351,43 @@ async def test_can_key_call_model_all_team_models_no_team_id_is_denied():
|
|||
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_can_team_access_model_all_team_models_expands_router_models():
|
||||
from litellm import Router
|
||||
from litellm.proxy._types import SpecialModelNames
|
||||
from litellm.proxy.auth.auth_checks import can_team_access_model
|
||||
|
||||
team_object = LiteLLM_TeamTable(
|
||||
team_id="team-123",
|
||||
models=[SpecialModelNames.all_team_models.value],
|
||||
)
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "allowed-model",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert (
|
||||
await can_team_access_model(
|
||||
model="allowed-model",
|
||||
team_object=team_object,
|
||||
llm_router=router,
|
||||
)
|
||||
is True
|
||||
)
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await can_team_access_model(
|
||||
model="blocked-model",
|
||||
team_object=team_object,
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_should_reconnect_once_on_db_connection_error():
|
||||
mock_prisma_client = MagicMock()
|
||||
|
|
|
|||
|
|
@ -565,6 +565,27 @@ def test_get_key_models_all_team_models_recursive_team():
|
|||
assert set(result) == {"model-a", "model-b"}
|
||||
|
||||
|
||||
def test_get_key_models_all_team_models_keeps_mixed_team_entries():
|
||||
from litellm.proxy.auth.model_checks import get_key_models
|
||||
from litellm.proxy._types import SpecialModelNames
|
||||
|
||||
user_api_key_dict = type(
|
||||
"obj",
|
||||
(object,),
|
||||
{
|
||||
"models": [SpecialModelNames.all_team_models.value],
|
||||
"team_id": "team-1",
|
||||
"team_models": [
|
||||
SpecialModelNames.all_team_models.value,
|
||||
"restricted-model",
|
||||
],
|
||||
},
|
||||
)()
|
||||
result = get_key_models(user_api_key_dict, ["model-a", "model-b"], {})
|
||||
assert SpecialModelNames.all_team_models.value not in result
|
||||
assert set(result) == {"model-a", "model-b", "restricted-model"}
|
||||
|
||||
|
||||
def test_get_team_models_all_team_models_expands():
|
||||
"""GH#30619: all-team-models in team_models should expand."""
|
||||
from litellm.proxy.auth.model_checks import get_team_models
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue