fix(proxy): gate v1 team filter and honor key allowlists

Only apply get_all_team_and_direct_access_models for admin or user-bound
keys, then intersect with key/team model restrictions to avoid empty lists
for service tokens and metadata leaks for restricted keys.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-06-10 17:02:33 +05:30
parent 72a960173a
commit 99c58a944f
No known key found for this signature in database
2 changed files with 169 additions and 33 deletions

View file

@ -12405,6 +12405,61 @@ def _deployment_matches_allowed_model_names(
)
def _get_v1_model_info_allowed_model_names(
user_api_key_dict: UserAPIKeyAuth,
llm_router: Router,
) -> Optional[Set[str]]:
"""Return key/team allowlisted public model names, or None if unrestricted."""
model_access_groups = llm_router.get_model_access_groups()
proxy_model_list = llm_router.get_model_names()
key_models = get_key_models(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
team_models = get_team_models(
team_models=user_api_key_dict.team_models,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
if not key_models and not team_models:
return None
return set(
get_complete_model_list(
key_models=key_models,
team_models=team_models,
proxy_model_list=proxy_model_list,
user_model=user_model,
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),
llm_router=llm_router,
return_wildcard_routes=False,
)
)
def _filter_v1_model_info_deployments(
all_models: List[dict],
allowed_model_names: Optional[Set[str]],
) -> List[dict]:
if allowed_model_names is None:
return all_models
return [
model
for model in all_models
if _deployment_matches_allowed_model_names(model, allowed_model_names)
]
def _should_apply_v1_team_access_filter(
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
"""Team membership filtering requires a resolvable user or admin role."""
return (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_id is not None
)
def _translate_model_name_for_response(model: dict) -> dict:
"""For team-scoped DB rows, replace `model_name` with the public name
in `model_info.team_public_model_name` before returning. The DB column
@ -12579,46 +12634,25 @@ async def model_info_v1( # noqa: PLR0915
# use internal routing keys (model_name_{team_id}_{uuid}) and were omitted
# when v1 resolved models only via public model_name strings.
all_models: List[dict] = copy.deepcopy(llm_router.model_list)
allowed_model_names = _get_v1_model_info_allowed_model_names(
user_api_key_dict=user_api_key_dict,
llm_router=llm_router,
)
if prisma_client is not None:
if prisma_client is not None and _should_apply_v1_team_access_filter(
user_api_key_dict=user_api_key_dict
):
all_models = await get_all_team_and_direct_access_models(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
llm_router=llm_router,
all_models=all_models,
)
else:
model_access_groups = llm_router.get_model_access_groups()
proxy_model_list = llm_router.get_model_names()
key_models = get_key_models(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
team_models = get_team_models(
team_models=user_api_key_dict.team_models,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
)
if key_models or team_models:
allowed_model_names = set(
get_complete_model_list(
key_models=key_models,
team_models=team_models,
proxy_model_list=proxy_model_list,
user_model=user_model,
infer_model_from_keys=general_settings.get(
"infer_model_from_keys", False
),
llm_router=llm_router,
return_wildcard_routes=False,
)
)
all_models = [
model
for model in all_models
if _deployment_matches_allowed_model_names(model, allowed_model_names)
]
all_models = _filter_v1_model_info_deployments(
all_models=all_models,
allowed_model_names=allowed_model_names,
)
all_models = [
_translate_model_name_for_response(

View file

@ -212,3 +212,105 @@ async def test_model_info_v1_list_path_translates_team_model_name(monkeypatch):
m for m in resp["data"] if m["model_name"] == "team-claude-sonnet"
)
assert team_model["model_info"]["access_via_team_ids"] == ["team-abc-123"]
@pytest.mark.asyncio
async def test_model_info_v1_no_user_id_with_db_skips_team_access_filter(monkeypatch):
"""Service/CI keys without user_id must not hit the team-membership filter."""
deployment = {
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4"},
"model_info": {"id": "global-id-1", "db_model": False},
}
router = MagicMock()
router.model_list = [deployment]
router.get_model_names.return_value = ["gpt-4"]
router.get_model_access_groups.return_value = {}
get_team_access = AsyncMock()
monkeypatch.setattr(ps, "user_model", None)
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
monkeypatch.setattr(ps, "llm_router", router)
monkeypatch.setattr(ps, "prisma_client", MagicMock())
monkeypatch.setattr(ps, "get_all_team_and_direct_access_models", get_team_access)
monkeypatch.setattr(
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
)
import litellm.proxy.agent_endpoints.model_list_helpers as mlh
monkeypatch.setattr(
mlh,
"append_agents_to_model_info",
AsyncMock(side_effect=lambda models, **kw: models),
)
caller = UserAPIKeyAuth(
user_id=None,
user_role=LitellmUserRoles.INTERNAL_USER,
models=[],
team_models=[],
)
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
assert [m["model_name"] for m in resp["data"]] == ["gpt-4"]
get_team_access.assert_not_called()
@pytest.mark.asyncio
async def test_model_info_v1_restricted_key_filters_after_team_enrichment(monkeypatch):
"""Key-level model allowlists must apply after DB team-access enrichment."""
team_row = _team_row()
global_row = {
"model_name": "gpt-4",
"litellm_params": {"model": "gpt-4"},
"model_info": {"id": "global-id-1", "db_model": False},
}
router = MagicMock()
router.model_list = [team_row, global_row]
router.get_model_names.return_value = ["gpt-4", "team-claude-sonnet"]
router.get_model_access_groups.return_value = {}
monkeypatch.setattr(ps, "user_model", None)
monkeypatch.setattr(ps, "llm_model_list", router.model_list)
monkeypatch.setattr(ps, "llm_router", router)
monkeypatch.setattr(ps, "prisma_client", MagicMock())
monkeypatch.setattr(
ps,
"get_all_team_and_direct_access_models",
AsyncMock(
side_effect=lambda all_models, **kwargs: [
{
**m,
"model_info": {
**m.get("model_info", {}),
**(
{"access_via_team_ids": ["team-abc-123"]}
if m.get("model_info", {}).get("team_id")
else {"direct_access": True}
),
},
}
for m in all_models
]
),
)
monkeypatch.setattr(
ps, "_enrich_model_info_with_litellm_data", lambda model, **kw: model
)
import litellm.proxy.agent_endpoints.model_list_helpers as mlh
monkeypatch.setattr(
mlh,
"append_agents_to_model_info",
AsyncMock(side_effect=lambda models, **kw: models),
)
caller = UserAPIKeyAuth(
user_id="user-1",
user_role=LitellmUserRoles.INTERNAL_USER,
models=["gpt-4"],
team_models=[],
)
resp = await ps.model_info_v1(user_api_key_dict=caller, litellm_model_id=None)
assert [m["model_name"] for m in resp["data"]] == ["gpt-4"]