fix(proxy): resolve entity access groups in the model listing endpoints (#36230)

* fix(proxy): resolve entity access groups in the model listing endpoints

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(proxy): reuse the fetched team object when listing models

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): cover key-level access group resolution in model listing

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: shivam <shivam@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-08-07 17:45:30 -07:00 committed by GitHub
parent e50a42051c
commit 1a45bf9afe
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 262 additions and 28 deletions

View file

@ -13,6 +13,7 @@ import asyncio
import math
import re
import time
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from fastapi import HTTPException, Request, status
@ -2834,7 +2835,7 @@ async def get_org_object(
async def _get_resources_from_access_groups(
access_group_ids: list[str],
access_group_ids: Sequence[str],
resource_field: Literal["access_model_names", "access_mcp_server_ids", "access_agent_ids"],
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
@ -2893,7 +2894,7 @@ async def _get_resources_from_access_groups(
async def _get_models_from_access_groups(
access_group_ids: list[str],
access_group_ids: Sequence[str],
prisma_client: PrismaClient | None = None,
user_api_key_cache: UserApiKeyCache | None = None,
proxy_logging_obj: ProxyLogging | None = None,

View file

@ -1,6 +1,7 @@
# What is this?
## Common checks for /v1/models and `/model/info`
import copy
from collections.abc import Sequence
from typing import Any, Final
import litellm
@ -178,8 +179,8 @@ def get_team_models(
def get_complete_model_list(
key_models: list[str],
team_models: list[str],
key_models: Sequence[str],
team_models: Sequence[str],
proxy_model_list: list[str],
user_model: str | None,
infer_model_from_keys: bool | None,
@ -203,7 +204,7 @@ def get_complete_model_list(
def append_unique(models):
for model in models:
if model not in unique_models:
if model not in unique_models and model != SpecialModelNames.no_default_models.value:
unique_models.append(model)
if key_models:

View file

@ -165,6 +165,7 @@ if TYPE_CHECKING:
from prisma.client import TransactionManager
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.models.team import LiteLLM_TeamTableCachedObj
from litellm.proxy.db.autorouter_session_rollup import AutoRouterTurnTransaction
from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction
@ -6419,6 +6420,74 @@ def construct_database_url_from_env_vars() -> str | None:
return None
async def _get_validated_team_object(
user_api_key_dict: "UserAPIKeyAuth",
team_id: str,
prisma_client: "PrismaClient",
user_api_key_cache: "UserApiKeyCache",
proxy_logging_obj: "ProxyLogging",
) -> "LiteLLM_TeamTableCachedObj":
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.management_endpoints.team_endpoints import validate_membership
team_object: Final = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_object)
return team_object
async def _get_team_object_for_access_groups(
team_id: str | None,
prisma_client: Optional["PrismaClient"],
user_api_key_cache: Optional["UserApiKeyCache"],
proxy_logging_obj: Optional["ProxyLogging"],
) -> Optional["LiteLLM_TeamTableCachedObj"]:
from litellm.proxy.auth.auth_checks import get_team_object
if team_id is None or prisma_client is None or user_api_key_cache is None or proxy_logging_obj is None:
return None
try:
return await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
except HTTPException:
verbose_proxy_logger.debug("Could not fetch team %s while listing models", team_id)
return None
async def _get_access_group_models(
user_api_key_dict: "UserAPIKeyAuth",
team_object: Optional["LiteLLM_TeamTableCachedObj"],
prisma_client: Optional["PrismaClient"],
user_api_key_cache: Optional["UserApiKeyCache"],
proxy_logging_obj: Optional["ProxyLogging"],
) -> tuple[str, ...]:
from litellm.proxy.auth.auth_checks import (
_get_models_from_access_groups,
get_authorized_resources_from_key_access_groups,
)
team_group_models: Final = await _get_models_from_access_groups(
access_group_ids=(team_object.access_group_ids or ()) if team_object is not None else (),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
key_group_models: Final = await get_authorized_resources_from_key_access_groups(
valid_token=user_api_key_dict,
team_object=team_object,
resource_field="access_model_names",
)
return tuple(dict.fromkeys((*team_group_models, *key_group_models)))
async def get_available_models_for_user(
user_api_key_dict: "UserAPIKeyAuth",
llm_router: Optional["Router"],
@ -6450,13 +6519,11 @@ async def get_available_models_for_user(
Returns:
List of model names available to the user
"""
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.auth.model_checks import (
get_complete_model_list,
get_key_models,
get_team_models,
)
from litellm.proxy.management_endpoints.team_endpoints import validate_membership
# Get proxy model list and access groups
if llm_router is None:
@ -6466,31 +6533,33 @@ async def get_available_models_for_user(
proxy_model_list = llm_router.get_model_names()
model_access_groups = llm_router.get_model_access_groups()
# Get key models
key_models = get_key_models(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
)
# Get team models
team_models: list[str] = user_api_key_dict.team_models
# If specific team_id is provided, validate and get team models
if team_id and prisma_client and proxy_logging_obj and user_api_key_cache:
key_models = []
team_object: Final = await get_team_object(
requested_team_object: Final = (
await _get_validated_team_object(
user_api_key_dict=user_api_key_dict,
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await validate_membership(user_api_key_dict=user_api_key_dict, team_table=team_object)
team_models = team_object.models
if team_id and prisma_client and proxy_logging_obj and user_api_key_cache
else None
)
team_models = get_team_models(
team_models=team_models,
key_models: Final[Sequence[str]] = (
()
if requested_team_object is not None
else get_key_models(
user_api_key_dict=user_api_key_dict,
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
)
)
team_models: Final = get_team_models(
team_models=(
requested_team_object.models if requested_team_object is not None else user_api_key_dict.team_models
),
proxy_model_list=proxy_model_list,
model_access_groups=model_access_groups,
include_model_access_groups=include_model_access_groups,
@ -6498,10 +6567,31 @@ async def get_available_models_for_user(
effective_team_id: Final = team_id or user_api_key_dict.team_id
access_group_models: Final = (
await _get_access_group_models(
user_api_key_dict=user_api_key_dict,
team_object=requested_team_object
or await _get_team_object_for_access_groups(
team_id=effective_team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
if key_models or team_models
else ()
)
granted_key_models: Final = (*key_models, *access_group_models) if key_models else key_models
granted_team_models: Final = (*team_models, *access_group_models) if team_models else team_models
# Get complete model list
all_models: Final = get_complete_model_list(
key_models=key_models,
team_models=team_models,
key_models=granted_key_models,
team_models=granted_team_models,
proxy_model_list=proxy_model_list,
user_model=user_model,
infer_model_from_keys=general_settings.get("infer_model_from_keys", False),

View file

@ -735,3 +735,28 @@ def test_add_known_models_refreshes_models_by_provider_for_wildcard_expansion():
litellm.vertex_language_models.discard(fake_model)
litellm.add_known_models(model_cost_map={})
assert fake_model not in litellm.models_by_provider["vertex_ai"]
def test_get_complete_model_list_drops_no_default_models_sentinel():
from litellm.proxy.auth.model_checks import get_complete_model_list
result = get_complete_model_list(
key_models=["no-default-models", "model-a"],
team_models=[],
proxy_model_list=["model-a", "model-b"],
user_model=None,
infer_model_from_keys=False,
)
assert result == ["model-a"]
def test_get_complete_model_list_sentinel_only_grants_nothing():
from litellm.proxy.auth.model_checks import get_complete_model_list
result = get_complete_model_list(
key_models=["no-default-models"],
team_models=["no-default-models"],
proxy_model_list=["model-a", "model-b"],
user_model=None,
infer_model_from_keys=False,
)
assert result == []

View file

@ -9,6 +9,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import (
create_model_info_response,
get_available_models_for_user,
hash_token,
is_known_model,
is_known_vector_store_index,
model_dump_with_preserved_fields,
@ -404,3 +405,119 @@ async def test_get_available_models_for_user_error_path_complete_list_raises(
general_settings={},
user_model=None,
)
@pytest.mark.asyncio
async def test_get_available_models_for_user_resolves_team_access_group_models(
monkeypatch,
):
from litellm.models.access_group import LiteLLM_AccessGroupTable
from litellm.models.team import LiteLLM_TeamTableCachedObj
team = LiteLLM_TeamTableCachedObj(
team_id="team-1",
models=["no-default-models"],
access_group_ids=["ag-1"],
)
access_group = LiteLLM_AccessGroupTable(
access_group_id="ag-1",
access_group_name="repro-group",
access_model_names=["model-a", "model-b"],
assigned_team_ids=["team-1"],
)
async def _get_team_object(**_kwargs):
return team
async def _get_access_object(**_kwargs):
return access_group
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object)
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_access_object", _get_access_object)
result = await get_available_models_for_user(
user_api_key_dict=UserAPIKeyAuth(
api_key="sk-test-key",
user_id="user-1",
team_id="team-1",
models=["all-team-models"],
team_models=["no-default-models"],
),
llm_router=_router_with_models(["model-a", "model-b", "model-c"]),
general_settings={},
user_model=None,
prisma_client=MagicMock(),
proxy_logging_obj=MagicMock(),
user_api_key_cache=MagicMock(),
)
assert sorted(result) == ["model-a", "model-b"]
@pytest.mark.asyncio
async def test_get_available_models_for_user_without_access_groups_grants_nothing(
monkeypatch,
):
from litellm.models.team import LiteLLM_TeamTableCachedObj
async def _get_team_object(**_kwargs):
return LiteLLM_TeamTableCachedObj(team_id="team-1", models=["no-default-models"])
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object)
result = await get_available_models_for_user(
user_api_key_dict=UserAPIKeyAuth(
api_key="sk-test-key",
user_id="user-1",
team_id="team-1",
models=["all-team-models"],
team_models=["no-default-models"],
),
llm_router=_router_with_models(["model-a", "model-b"]),
general_settings={},
user_model=None,
prisma_client=MagicMock(),
proxy_logging_obj=MagicMock(),
user_api_key_cache=MagicMock(),
)
assert result == []
@pytest.mark.asyncio
async def test_get_available_models_for_user_resolves_key_access_group_models(
monkeypatch,
):
from litellm.models.access_group import LiteLLM_AccessGroupTable
from litellm.models.team import LiteLLM_TeamTableCachedObj
async def _get_team_object(**_kwargs):
return LiteLLM_TeamTableCachedObj(team_id="team-1", models=["no-default-models"])
async def _get_access_object(**_kwargs):
return LiteLLM_AccessGroupTable(
access_group_id="ag-1",
access_group_name="key-group",
access_model_names=["model-b"],
assigned_key_ids=[hash_token("sk-test-key")],
)
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_object", _get_team_object)
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_access_object", _get_access_object)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock())
result = await get_available_models_for_user(
user_api_key_dict=UserAPIKeyAuth(
api_key="sk-test-key",
user_id="user-1",
team_id="team-1",
models=["no-default-models"],
team_models=["no-default-models"],
access_group_ids=["ag-1"],
),
llm_router=_router_with_models(["model-a", "model-b"]),
general_settings={},
user_model=None,
prisma_client=MagicMock(),
proxy_logging_obj=MagicMock(),
user_api_key_cache=MagicMock(),
)
assert result == ["model-b"]