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:
milan 2026-08-04 23:27:32 +00:00
parent e2950a8995
commit 8390e6d966
3 changed files with 275 additions and 4 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, 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,

View file

@ -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(

View file

@ -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