mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge 1b8d18016e into 9fd78ff6f4
This commit is contained in:
commit
0ec9971108
2 changed files with 95 additions and 1 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue