From 8390e6d9664f98365b575d100455b1b2e357f0b4 Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 4 Aug 2026 23:27:32 +0000 Subject: [PATCH] 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> --- litellm/proxy/auth/auth_checks.py | 5 +- litellm/proxy/utils.py | 106 ++++++++++- .../proxy/utils/helpers/test_model_access.py | 168 ++++++++++++++++++ 3 files changed, 275 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 5ef4eb471ad..3fddcf542dd 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 99e566c2da1..7e28512e0b6 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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( diff --git a/tests/test_litellm/proxy/utils/helpers/test_model_access.py b/tests/test_litellm/proxy/utils/helpers/test_model_access.py index 59268e1427b..2253c3c3ed6 100644 --- a/tests/test_litellm/proxy/utils/helpers/test_model_access.py +++ b/tests/test_litellm/proxy/utils/helpers/test_model_access.py @@ -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