fix: include model_access_groups when expanding all-team-models in get_team_models (#30622)

This commit is contained in:
Zang Peiyu 2026-06-22 20:43:20 +08:00 committed by GitHub
parent 8f4a2bccf0
commit ad2b395c9d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 67 additions and 3 deletions

View file

@ -122,9 +122,14 @@ def get_key_models(
SpecialModelNames.all_team_models.value in all_models
and user_api_key_dict.team_id is not None
):
all_models = list(
user_api_key_dict.team_models
) # copy to avoid mutating cached objects
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)
if include_model_access_groups:
all_models.extend(model_access_groups.keys())
if SpecialModelNames.all_proxy_models.value in all_models:
all_models = list(proxy_model_list) # copy to avoid mutating caller's list
if include_model_access_groups:
@ -160,6 +165,12 @@ def get_team_models(
all_models_set.update(team_models)
if SpecialModelNames.all_team_models.value in all_models_set:
all_models_set.update(team_models)
# GH#30619: expand all-team-models sentinel
# to the actual proxy model list
all_models_set.discard(SpecialModelNames.all_team_models.value)
all_models_set.update(proxy_model_list)
if include_model_access_groups:
all_models_set.update(model_access_groups.keys())
if SpecialModelNames.all_proxy_models.value in all_models_set:
all_models_set.update(proxy_model_list)
if include_model_access_groups:

View file

@ -543,3 +543,56 @@ async def test_get_available_models_for_user_expands_query_team_wildcard(
)
assert "openai/gpt-4o-mini" in result
def test_get_key_models_all_team_models_recursive_team():
"""GH#30619: when key and team both have all-team-models,
the sentinel should expand to proxy_model_list."""
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],
},
)()
proxy_model_list = ["model-a", "model-b"]
result = get_key_models(user_api_key_dict, proxy_model_list, {})
assert SpecialModelNames.all_team_models.value not in result
assert set(result) == {"model-a", "model-b"}
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
from litellm.proxy._types import SpecialModelNames
result = get_team_models(
[SpecialModelNames.all_team_models.value],
["model-a", "model-b"],
{},
)
assert SpecialModelNames.all_team_models.value not in result
assert set(result) == {"model-a", "model-b"}
def test_get_team_models_all_team_models_expands_with_access_groups():
"""GH#30619: all-team-models with include_model_access_groups
should include access group keys."""
from litellm.proxy.auth.model_checks import get_team_models
from litellm.proxy._types import SpecialModelNames
result = get_team_models(
[SpecialModelNames.all_team_models.value],
["model-a", "model-b"],
{"group-1": ["g1-model"], "group-2": ["g2-model"]},
include_model_access_groups=True,
)
assert SpecialModelNames.all_team_models.value not in result
assert "model-a" in result
assert "model-b" in result
assert "group-1" in result
assert "group-2" in result