mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
Merge pull request #41094 from BerriAI/litellm_model_group_info_proxy_admin_all_models
fix(proxy): show all model groups to proxy admins in /model_group/info
This commit is contained in:
commit
09a188b583
2 changed files with 106 additions and 23 deletions
|
|
@ -15557,18 +15557,34 @@ async def model_group_info(
|
|||
from litellm.proxy.utils import get_available_models_for_user
|
||||
|
||||
# Get available models for the user
|
||||
all_models_str: Final = 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,
|
||||
team_id=None,
|
||||
include_model_access_groups=False,
|
||||
only_model_access_groups=False,
|
||||
return_wildcard_routes=False,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
is_proxy_admin: Final = user_api_key_dict.user_role in (
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
)
|
||||
all_models_str: Final = (
|
||||
get_complete_model_list(
|
||||
key_models=(),
|
||||
team_models=(),
|
||||
proxy_model_list=llm_router.get_model_names(),
|
||||
user_model=user_model,
|
||||
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),
|
||||
return_wildcard_routes=False,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
if is_proxy_admin
|
||||
else 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,
|
||||
team_id=None,
|
||||
include_model_access_groups=False,
|
||||
only_model_access_groups=False,
|
||||
return_wildcard_routes=False,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
)
|
||||
model_groups: list[ModelGroupInfoProxy] = _get_model_group_info(
|
||||
llm_router=llm_router, all_models_str=all_models_str, model_group=model_group
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ Pins (PR2):
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -128,7 +128,6 @@ def test_v1_model_info_no_model_list_error(client, auth_as, null_router, path):
|
|||
assert "LLM Model List not loaded" in response.text
|
||||
|
||||
|
||||
|
||||
def test_get_proxy_model_info_surfaces_supports_parallel_function_calling(local_model_cost_map):
|
||||
"""``GET /v1/model/info`` enriches each deployment through ``_get_proxy_model_info``; a registry
|
||||
entry declaring parallel function calling must land in ``model_info`` instead of null."""
|
||||
|
|
@ -161,9 +160,7 @@ def test_v1_model_info_star_wildcard_filter_keeps_provider_expansion(monkeypatch
|
|||
router.get_model_list = MagicMock(return_value=[deployment])
|
||||
monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models)
|
||||
|
||||
expanded_deployments = proxy_server.expand_wildcard_deployments_for_model_info(
|
||||
[deployment]
|
||||
)
|
||||
expanded_deployments = proxy_server.expand_wildcard_deployments_for_model_info([deployment])
|
||||
allowed_model_names = proxy_server._get_v1_model_info_allowed_model_names(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="sk-test",
|
||||
|
|
@ -308,6 +305,80 @@ def test_model_group_info_invalid_method(client, auth_as, null_router):
|
|||
assert len(response.content) > 0
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model_group_info_router(monkeypatch):
|
||||
from litellm.types.proxy.management_endpoints.model_management_endpoints import ModelGroupInfoProxy
|
||||
|
||||
model_names = ["gpt-4", "claude-3"]
|
||||
router = MagicMock()
|
||||
router.get_model_names.return_value = model_names
|
||||
router.get_model_access_groups.return_value = {}
|
||||
router.get_model_list.return_value = []
|
||||
|
||||
def model_group_info(*, llm_router, all_models_str, model_group):
|
||||
return [ModelGroupInfoProxy(model_group=name, providers=[]) for name in all_models_str]
|
||||
|
||||
async def append_agents_to_model_group(*, model_groups, user_api_key_dict):
|
||||
return model_groups
|
||||
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(proxy_server, "llm_model_list", [{"model_name": name} for name in model_names])
|
||||
monkeypatch.setattr(proxy_server, "user_model", None)
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", None)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", None)
|
||||
monkeypatch.setattr(proxy_server, "_get_model_group_info", model_group_info)
|
||||
|
||||
from litellm.proxy.agent_endpoints import model_list_helpers
|
||||
|
||||
monkeypatch.setattr(
|
||||
model_list_helpers,
|
||||
"append_agents_to_model_group",
|
||||
AsyncMock(side_effect=append_agents_to_model_group),
|
||||
)
|
||||
return router
|
||||
|
||||
|
||||
@pytest.mark.parametrize("admin_role", ["proxy_admin", "proxy_admin_viewer"])
|
||||
def test_model_group_info_proxy_admin_ignores_key_model_restriction(
|
||||
client, auth_as, model_group_info_router, admin_role
|
||||
):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
with auth_as(LitellmUserRoles(admin_role), models=["no-default-models"]):
|
||||
response = client.get("/model_group/info")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [model["model_group"] for model in response.json()["data"]] == ["gpt-4", "claude-3"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("admin_role", ["proxy_admin", "proxy_admin_viewer"])
|
||||
def test_model_group_info_proxy_admin_expands_wildcard_deployments(client, auth_as, model_group_info_router, admin_role):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard
|
||||
|
||||
model_group_info_router.get_model_names.return_value = ["gpt-4", "anthropic/*"]
|
||||
known_anthropic_models = get_known_models_from_wildcard(wildcard_model="anthropic/*")
|
||||
assert known_anthropic_models
|
||||
|
||||
with auth_as(LitellmUserRoles(admin_role), models=["no-default-models"]):
|
||||
response = client.get("/model_group/info")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [model["model_group"] for model in response.json()["data"]] == ["gpt-4", *known_anthropic_models]
|
||||
|
||||
|
||||
def test_model_group_info_internal_user_key_model_restriction_applies(client, auth_as, model_group_info_router):
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
with auth_as(LitellmUserRoles.INTERNAL_USER, models=["gpt-4"]):
|
||||
response = client.get("/model_group/info")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [model["model_group"] for model in response.json()["data"]] == ["gpt-4"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GET /v2/model/info?exclude_auto_routers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -399,14 +470,10 @@ def test_v2_model_info_exclude_auto_routers_shrinks_total_count(client, auth_as,
|
|||
assert len(payload["data"]) == payload["total_count"]
|
||||
|
||||
|
||||
def test_v2_model_info_exclude_auto_routers_paginates_over_the_filtered_set(
|
||||
client, auth_as, mixed_auto_router_router
|
||||
):
|
||||
def test_v2_model_info_exclude_auto_routers_paginates_over_the_filtered_set(client, auth_as, mixed_auto_router_router):
|
||||
"""Page size applies to the filtered list, so no page silently comes back short."""
|
||||
with auth_as():
|
||||
response = client.get(
|
||||
"/v2/model/info", params={"exclude_auto_routers": "true", "page": 1, "size": 1}
|
||||
)
|
||||
response = client.get("/v2/model/info", params={"exclude_auto_routers": "true", "page": 1, "size": 1})
|
||||
payload = response.json()
|
||||
assert payload["total_count"] == 2
|
||||
assert payload["total_pages"] == 2
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue