diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 57566bed406..a29f7077951 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10998,17 +10998,12 @@ 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 - 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", []) - ] + 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 @@ -12455,19 +12450,14 @@ def _filter_v1_model_info_deployments( ] -async def _should_apply_v1_team_access_filter( +def _should_apply_v1_team_access_filter( user_api_key_dict: UserAPIKeyAuth, - prisma_client: PrismaClient, ) -> bool: - """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} + """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 ) - return user_db_object is not None def _translate_model_name_for_response(model: dict) -> dict: @@ -12649,9 +12639,8 @@ async def model_info_v1( # noqa: PLR0915 llm_router=llm_router, ) - 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, + 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, 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 bc9a44010fa..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 @@ -270,12 +270,6 @@ 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) @@ -320,87 +314,3 @@ 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