From 99c58a944fe4b7a4830dcbe731bbd7993ac77cef Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 10 Jun 2026 17:02:33 +0530 Subject: [PATCH] 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 --- litellm/proxy/proxy_server.py | 100 +++++++++++------ .../test_team_model_name_translation.py | 102 ++++++++++++++++++ 2 files changed, 169 insertions(+), 33 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 69e9fcf8b83..a29f7077951 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py index 1d46113ed3c..df9b45fd29d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py +++ b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py @@ -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"]