diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a29f7077951..57566bed406 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10998,12 +10998,17 @@ async def get_all_team_and_direct_access_models( _model["model_info"]["direct_access"] = True ## FILTER OUT MODELS THAT ARE NOT IN DIRECT_ACCESS_MODELS OR ACCESS_VIA_TEAM_IDS - only show user models they can call - all_models = [ - _model - for _model in all_models - if _model.get("model_info", {}).get("direct_access", False) - or _model.get("model_info", {}).get("access_via_team_ids", []) - ] + should_filter_by_user_access = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or user_teams is not None + ) + if should_filter_by_user_access: + all_models = [ + _model + for _model in all_models + if _model.get("model_info", {}).get("direct_access", False) + or _model.get("model_info", {}).get("access_via_team_ids", []) + ] return all_models @@ -12450,14 +12455,19 @@ def _filter_v1_model_info_deployments( ] -def _should_apply_v1_team_access_filter( +async def _should_apply_v1_team_access_filter( user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient, ) -> 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 + """Team membership filtering requires admin role or a DB-backed user row.""" + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: + return True + if user_api_key_dict.user_id is None: + return False + user_db_object = await UserRepository(prisma_client).table.find_unique( + where={"user_id": user_api_key_dict.user_id} ) + return user_db_object is not None def _translate_model_name_for_response(model: dict) -> dict: @@ -12639,8 +12649,9 @@ async def model_info_v1( # noqa: PLR0915 llm_router=llm_router, ) - if prisma_client is not None and _should_apply_v1_team_access_filter( - user_api_key_dict=user_api_key_dict + if prisma_client is not None and await _should_apply_v1_team_access_filter( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, ): all_models = await get_all_team_and_direct_access_models( user_api_key_dict=user_api_key_dict, 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 df9b45fd29d..bc9a44010fa 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 @@ -270,6 +270,12 @@ async def test_model_info_v1_restricted_key_filters_after_team_enrichment(monkey router.get_model_names.return_value = ["gpt-4", "team-claude-sonnet"] router.get_model_access_groups.return_value = {} + class MockUserRepository: + def __init__(self, prisma_client): + self.table = MagicMock() + self.table.find_unique = AsyncMock(return_value=MagicMock(model_dump=lambda: {"user_id": "user-1", "teams": []})) + + monkeypatch.setattr(ps, "UserRepository", MockUserRepository) monkeypatch.setattr(ps, "user_model", None) monkeypatch.setattr(ps, "llm_model_list", router.model_list) monkeypatch.setattr(ps, "llm_router", router) @@ -314,3 +320,87 @@ async def test_model_info_v1_restricted_key_filters_after_team_enrichment(monkey 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"] + + +@pytest.mark.asyncio +async def test_model_info_v1_missing_db_user_returns_deployments(monkeypatch): + """Keys with user_id but no DB row must not lose their model list.""" + 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 = {} + + class MockUserRepository: + def __init__(self, prisma_client): + self.table = MagicMock() + self.table.find_unique = AsyncMock(return_value=None) + + get_team_access = AsyncMock() + monkeypatch.setattr(ps, "UserRepository", MockUserRepository) + 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="missing-user", + 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_get_all_team_and_direct_access_models_missing_user_skips_filter( + monkeypatch, +): + """Defense in depth: unresolved user_id must not empty the model list.""" + models = [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "global-id-1"}, + } + ] + + class MockUserRepository: + def __init__(self, prisma_client): + self.table = MagicMock() + self.table.find_unique = AsyncMock(return_value=None) + + monkeypatch.setattr(ps, "UserRepository", MockUserRepository) + + caller = UserAPIKeyAuth( + user_id="missing-user", + user_role=LitellmUserRoles.INTERNAL_USER, + models=[], + team_models=[], + ) + result = await ps.get_all_team_and_direct_access_models( + user_api_key_dict=caller, + prisma_client=MagicMock(), + llm_router=MagicMock(), + all_models=models, + ) + + assert result == models