mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
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:
parent
e50a42051c
commit
1a45bf9afe
5 changed files with 262 additions and 28 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue