mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): include access group models in /v1/models for teams and keys
Models granted to a team or key through a unified access group were callable but absent from the model listing, because get_available_models_for_user only read team.models / key.models while request-time auth falls back to access_group_ids. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e2950a8995
commit
8390e6d966
3 changed files with 275 additions and 4 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, Union, cast
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
|
|
@ -2836,7 +2837,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,
|
||||
|
|
@ -2895,7 +2896,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,
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from dataclasses import dataclass, field
|
|||
from datetime import date, datetime, timedelta, timezone
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
from itertools import chain
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, Union, cast, overload
|
||||
|
||||
from litellm import _custom_logger_compatible_callbacks_literal
|
||||
|
|
@ -165,6 +166,7 @@ if TYPE_CHECKING:
|
|||
from prisma.client import TransactionManager
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
||||
from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction
|
||||
|
||||
Span = Union[_Span, Any]
|
||||
|
|
@ -6378,11 +6380,12 @@ async def get_available_models_for_user(
|
|||
|
||||
# Get team models
|
||||
team_models: list[str] = user_api_key_dict.team_models
|
||||
team_object: LiteLLM_TeamTableCachedObj | None = None
|
||||
|
||||
# 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(
|
||||
team_object = await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
|
|
@ -6415,7 +6418,106 @@ async def get_available_models_for_user(
|
|||
team_id=effective_team_id,
|
||||
)
|
||||
|
||||
return all_models
|
||||
if only_model_access_groups:
|
||||
return all_models
|
||||
|
||||
access_group_models: Final = await _get_models_from_unified_access_groups(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
team_object=team_object,
|
||||
effective_team_id=effective_team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if not access_group_models:
|
||||
return all_models
|
||||
|
||||
# mutable-ok: the documented return contract is list[str]; built in one shot, never mutated
|
||||
return list(dict.fromkeys(chain(all_models, access_group_models)))
|
||||
|
||||
|
||||
async def _get_team_access_group_ids(
|
||||
team_object: Optional["LiteLLM_TeamTableCachedObj"],
|
||||
effective_team_id: str | None,
|
||||
prisma_client: Optional["PrismaClient"],
|
||||
user_api_key_cache: Optional["UserApiKeyCache"],
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
) -> tuple[str, ...]:
|
||||
"""
|
||||
Access group ids assigned to the caller's team, using the already resolved
|
||||
team object when available and otherwise looking the team up (cache backed).
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
|
||||
if team_object is not None:
|
||||
return tuple(team_object.access_group_ids or ())
|
||||
|
||||
if not effective_team_id or prisma_client is None or user_api_key_cache is None:
|
||||
return ()
|
||||
|
||||
try:
|
||||
fetched_team: Final = await get_team_object(
|
||||
team_id=effective_team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
"Could not resolve team %s while listing access group models",
|
||||
effective_team_id,
|
||||
)
|
||||
return ()
|
||||
|
||||
return tuple(fetched_team.access_group_ids or ())
|
||||
|
||||
|
||||
async def _get_models_from_unified_access_groups(
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
team_object: Optional["LiteLLM_TeamTableCachedObj"],
|
||||
effective_team_id: str | None,
|
||||
prisma_client: Optional["PrismaClient"],
|
||||
user_api_key_cache: Optional["UserApiKeyCache"],
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
) -> tuple[str, ...]:
|
||||
"""
|
||||
Model names granted through unified access groups (LiteLLM_AccessGroupTable)
|
||||
assigned to the caller's team or key.
|
||||
|
||||
Mirrors the auth-time resolution in can_team_access_model and
|
||||
_key_access_group_grants_model, so a model the caller can successfully call
|
||||
is not missing from the listing.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_get_models_from_access_groups,
|
||||
get_authorized_resources_from_key_access_groups,
|
||||
)
|
||||
|
||||
team_access_group_ids: Final = await _get_team_access_group_ids(
|
||||
team_object=team_object,
|
||||
effective_team_id=effective_team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
team_group_models: Final = (
|
||||
await _get_models_from_access_groups(
|
||||
access_group_ids=team_access_group_ids,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if team_access_group_ids
|
||||
else ()
|
||||
)
|
||||
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(chain(team_group_models, key_group_models)))
|
||||
|
||||
|
||||
def create_model_info_response(
|
||||
|
|
|
|||
|
|
@ -404,3 +404,171 @@ async def test_get_available_models_for_user_error_path_complete_list_raises(
|
|||
general_settings={},
|
||||
user_model=None,
|
||||
)
|
||||
|
||||
|
||||
def _access_group_key(team_id="team-1", team_models=None, key_access_group_ids=None):
|
||||
return UserAPIKeyAuth(
|
||||
api_key="sk-test-key",
|
||||
user_id="user-1",
|
||||
team_id=team_id,
|
||||
models=[],
|
||||
team_models=team_models if team_models is not None else ["model-a"],
|
||||
access_group_ids=key_access_group_ids,
|
||||
)
|
||||
|
||||
|
||||
def _patch_team_access_groups(monkeypatch, access_group_ids, group_models):
|
||||
"""Stub the unified access group DB boundary used by the listing."""
|
||||
team_lookups = []
|
||||
resolved_for = []
|
||||
|
||||
async def _fake_get_team_object(**kwargs):
|
||||
team_lookups.append(kwargs["team_id"])
|
||||
team = MagicMock()
|
||||
team.access_group_ids = access_group_ids
|
||||
return team
|
||||
|
||||
async def _fake_models_from_access_groups(**kwargs):
|
||||
resolved_for.append(list(kwargs["access_group_ids"]))
|
||||
return list(group_models)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.auth_checks.get_team_object", _fake_get_team_object
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.auth_checks._get_models_from_access_groups",
|
||||
_fake_models_from_access_groups,
|
||||
)
|
||||
return team_lookups, resolved_for
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_available_models_for_user_includes_team_access_group_models(
|
||||
monkeypatch,
|
||||
):
|
||||
"""
|
||||
Regression: a team restricted to `model-a` that is assigned an access group
|
||||
granting `model-b` must see `model-b` in /v1/models. Before the fix the
|
||||
listing only consulted team.models, so `model-b` was missing even though the
|
||||
team key could call it.
|
||||
"""
|
||||
_team_lookups, resolved_for = _patch_team_access_groups(
|
||||
monkeypatch, access_group_ids=["ag-1"], group_models=["model-b"]
|
||||
)
|
||||
|
||||
result = await get_available_models_for_user(
|
||||
user_api_key_dict=_access_group_key(),
|
||||
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(),
|
||||
include_model_access_groups=True,
|
||||
)
|
||||
|
||||
assert {
|
||||
"result": sorted(result),
|
||||
"resolved_access_group_ids": resolved_for,
|
||||
} == {
|
||||
"result": ["model-a", "model-b"],
|
||||
"resolved_access_group_ids": [["ag-1"]],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_available_models_for_user_access_group_models_are_deduped(
|
||||
monkeypatch,
|
||||
):
|
||||
"""A model granted both directly and via an access group is listed once."""
|
||||
_patch_team_access_groups(
|
||||
monkeypatch, access_group_ids=["ag-1"], group_models=["model-a", "model-b"]
|
||||
)
|
||||
|
||||
result = await get_available_models_for_user(
|
||||
user_api_key_dict=_access_group_key(),
|
||||
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-a", "model-b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_available_models_for_user_includes_key_access_group_models(
|
||||
monkeypatch,
|
||||
):
|
||||
"""Models from an access group that authorizes the key are also listed."""
|
||||
_patch_team_access_groups(
|
||||
monkeypatch, access_group_ids=[], group_models=[]
|
||||
)
|
||||
|
||||
async def _fake_key_resources(**kwargs):
|
||||
assert kwargs["resource_field"] == "access_model_names"
|
||||
return ["model-b"]
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.auth_checks.get_authorized_resources_from_key_access_groups",
|
||||
_fake_key_resources,
|
||||
)
|
||||
|
||||
result = await get_available_models_for_user(
|
||||
user_api_key_dict=_access_group_key(key_access_group_ids=["ag-key"]),
|
||||
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 sorted(result) == ["model-a", "model-b"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_available_models_for_user_without_access_groups_is_unchanged(
|
||||
monkeypatch,
|
||||
):
|
||||
"""A team with no access groups still lists exactly its own models."""
|
||||
_patch_team_access_groups(
|
||||
monkeypatch, access_group_ids=[], group_models=["should-not-appear"]
|
||||
)
|
||||
|
||||
result = await get_available_models_for_user(
|
||||
user_api_key_dict=_access_group_key(),
|
||||
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-a"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_available_models_for_user_only_model_access_groups_skips_expansion(
|
||||
monkeypatch,
|
||||
):
|
||||
"""only_model_access_groups returns router access group names, not access group members."""
|
||||
_patch_team_access_groups(
|
||||
monkeypatch, access_group_ids=["ag-1"], group_models=["model-b"]
|
||||
)
|
||||
|
||||
result = await get_available_models_for_user(
|
||||
user_api_key_dict=_access_group_key(),
|
||||
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(),
|
||||
only_model_access_groups=True,
|
||||
)
|
||||
|
||||
assert "model-b" not in result
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue