diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ae559f30857..38c044fd5ce 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15815,6 +15815,42 @@ def _get_v1_model_info_allowed_model_names( ) +async def _get_v1_model_info_allowed_model_names_with_access_groups( + user_api_key_dict: UserAPIKeyAuth, + llm_router: Router, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache | None, + proxy_logging_obj: ProxyLogging | None, +) -> set[str] | None: # mutable-ok: preserve the existing model-filter helper contract + fallback_allowed_models: Final = _get_v1_model_info_allowed_model_names( + user_api_key_dict=user_api_key_dict, + llm_router=llm_router, + ) + if fallback_allowed_models is None: + return None + + from litellm.proxy.utils import get_available_models_for_user + + try: + return set( + await get_available_models_for_user( + user_api_key_dict=user_api_key_dict, + llm_router=llm_router, + general_settings=general_settings, + user_model=user_model, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + user_api_key_cache=user_api_key_cache, + ) + ) + except Exception: # noqa: BLE001 -- preserve the legacy listing when access-group enrichment fails + verbose_proxy_logger.debug( + "Could not resolve access group models for /v1/model/info", + exc_info=True, + ) + return fallback_allowed_models + + def _filter_v1_model_info_deployments( all_models: list[dict], allowed_model_names: set[str] | None, @@ -16020,9 +16056,12 @@ async def model_info_v1( all_models = expand_wildcard_deployments_for_model_info(all_models) - allowed_model_names: Final = _get_v1_model_info_allowed_model_names( + allowed_model_names: Final = await _get_v1_model_info_allowed_model_names_with_access_groups( user_api_key_dict=user_api_key_dict, llm_router=llm_router, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, ) all_models = _filter_v1_model_info_deployments( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index df05cf0987e..41342a71a06 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5860,6 +5860,61 @@ async def test_boot_warns_that_a_shadowed_database_value_will_never_apply(tmp_pa assert proxy_config.settings["max_parallel_requests"] == 7 +@pytest.mark.asyncio +async def test_get_v1_model_info_allowed_names_includes_inherited_team_access_group_models(monkeypatch): + from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth + + router = MagicMock() + router.get_model_names.return_value = ["model-a"] + router.get_model_access_groups.return_value = {} + available_models = AsyncMock(return_value=["model-a"]) + monkeypatch.setattr("litellm.proxy.utils.get_available_models_for_user", available_models) + + result = await proxy_server_module._get_v1_model_info_allowed_model_names_with_access_groups( + user_api_key_dict=UserAPIKeyAuth( + api_key="sk-test", + models=[SpecialModelNames.all_team_models.value], + team_id="team-1", + team_models=[SpecialModelNames.no_default_models.value], + ), + llm_router=router, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert result == {"model-a"} + available_models.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_get_v1_model_info_allowed_names_does_not_broaden_empty_access_group(monkeypatch): + from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth + + router = MagicMock() + router.get_model_names.return_value = ["model-a"] + router.get_model_access_groups.return_value = {} + monkeypatch.setattr( + "litellm.proxy.utils.get_available_models_for_user", + AsyncMock(return_value=[]), + ) + + result = await proxy_server_module._get_v1_model_info_allowed_model_names_with_access_groups( + user_api_key_dict=UserAPIKeyAuth( + api_key="sk-test", + models=[SpecialModelNames.all_team_models.value], + team_id="team-1", + team_models=[SpecialModelNames.no_default_models.value], + ), + llm_router=router, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert result == set() + + @pytest.mark.asyncio async def test_model_info_v1_oci_secrets_not_leaked(): """