From 9fb55fe829ed26a8407f4fe33e4c7c2b79425e61 Mon Sep 17 00:00:00 2001 From: songkuan-zheng <252822057+songkuan-zheng@users.noreply.github.com> Date: Tue, 19 May 2026 12:46:51 +0000 Subject: [PATCH] fix(proxy): apply user.models filter on /v1/models (BerriAI/litellm#26420) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit GET /v1/models was filtering by key.models and team.models but silently ignoring LiteLLM_UserTable.models (the "Personal Models" field). A user restricted to a subset of models would see the full proxy list on the discovery endpoint even though calling any restricted model returned 401 at /v1/chat/completions — an inconsistency between listing and inference enforcement. Root cause: get_available_models_for_user() in litellm/proxy/utils.py sources the model list from user_api_key_dict (key + team view) only. The user object's models field is loaded by user_api_key_auth but deliberately not spliced into UserAPIKeyAuth. Fix: add get_user_models() + filter_models_by_user_access() helpers in litellm/proxy/auth/model_checks.py mirroring the existing key/team helpers, then call them as a final intersection step inside get_available_models_for_user(). Defense-in-depth — the filter only narrows what the key/team layer already permitted, never expands it. Semantics match can_user_call_model() at inference time: - user.models=[] -> no filter - user.models=['no-default-models'] -> empty list - user.models=['all-proxy-models'] -> no filter - raw names / access groups -> exact + group expansion - wildcards ('anthropic/*', '*') -> fnmatch.fnmatchcase - user_id None (master/service key) -> no filter - get_user_object failure -> no filter (swallow, log) scope=expand admin branch in proxy_server.py untouched — intentional admin escape hatch gated by _user_has_admin_privileges. Tests (4 tiers, 28 new cases): - Unit (test_model_checks.py): get_user_models + filter helper, 11 cases - Async integration (test_proxy_utils.py): _apply_user_models_filter across 6 branches incl. user_id None, no-default-models, all-proxy, access groups, wildcard, exception swallow — 10 cases - Route-level FastAPI TestClient (discovery_endpoints/test_v1_models_user_filter.py): real model_list handler → real get_available_models_for_user → mocked get_user_object, 7 cases - E2E (e2e/cases/15): real proxy + Postgres, dynamic users via /user/new + /key/generate, verified PASS post-fix and FAIL pre-fix on the same live proxy. Wired into e2e/tools/run-all-cases. Closes internal mirror of BerriAI/litellm#26420. --- litellm/proxy/auth/model_checks.py | 65 +++++ litellm/proxy/utils.py | 90 +++++++ .../test_v1_models_user_filter.py | 240 ++++++++++++++++++ 3 files changed, 395 insertions(+) create mode 100644 tests/test_litellm/proxy/discovery_endpoints/test_v1_models_user_filter.py diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 5d5ab4f224f..a1aa4091d7e 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -1,6 +1,7 @@ # What is this? ## Common checks for /v1/models and `/model/info` import copy +import fnmatch from typing import Any, Dict, List, Optional, Set import litellm @@ -192,6 +193,70 @@ def get_team_models( return all_models +def get_user_models( + user_models: List[str], + proxy_model_list: List[str], + model_access_groups: Dict[str, List[str]], + include_model_access_groups: Optional[bool] = False, +) -> List[str]: + """ + Returns: + - List of model name strings allowed by `LiteLLM_UserTable.models` + (the "Personal Models" field). + - Empty list if no models set. + - Mirrors `get_team_models` semantics (sans `all-team-models`, + which has no meaning at the user scope). + + Used by `/v1/models` to apply per-user access restrictions on the + listing path so it stays consistent with `can_user_call_model` at + inference time (see BerriAI/litellm#26420). + """ + all_models_set: Set[str] = set() + if len(user_models) > 0: + all_models_set.update(user_models) + if SpecialModelNames.all_proxy_models.value in all_models_set: + all_models_set.update(proxy_model_list) + if include_model_access_groups: + all_models_set.update(model_access_groups.keys()) + + all_models = _get_models_from_access_groups( + model_access_groups=model_access_groups, + all_models=list(all_models_set), + include_model_access_groups=include_model_access_groups, + ) + + # deduplicate while preserving order + all_models = list(dict.fromkeys(all_models)) + + verbose_proxy_logger.debug("ALL USER MODELS - {}".format(len(all_models))) + return all_models + + +def filter_models_by_user_access( + models: List[str], + user_allowed_models: List[str], +) -> List[str]: + """ + Return the subset of `models` that the user is allowed to see, given + the (already-expanded) `user_allowed_models` list. Supports exact + match plus `fnmatch` wildcards (e.g. `anthropic/*`, `*`). + + Caller is responsible for short-circuiting before calling when + `user_allowed_models` is empty, contains `all-proxy-models` + (no filter), or contains `no-default-models` (empty result). + Order of `models` is preserved. + """ + exact = {m for m in user_allowed_models if "*" not in m} + patterns = [m for m in user_allowed_models if "*" in m] + out: List[str] = [] + for m in models: + if m in exact: + out.append(m) + elif patterns and any(fnmatch.fnmatchcase(m, p) for p in patterns): + out.append(m) + return out + + def get_complete_model_list( key_models: List[str], team_models: List[str], diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 781f0e9f301..8202367386f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -6437,9 +6437,99 @@ async def get_available_models_for_user( team_id=effective_team_id, ) + # Apply user-level (Personal Models) restriction so /v1/models is + # consistent with can_user_call_model at inference time + # (BerriAI/litellm#26420). Defense-in-depth: this only ever narrows + # the list, never widens it. + all_models = await _apply_user_models_filter( + all_models=all_models, + user_api_key_dict=user_api_key_dict, + proxy_model_list=proxy_model_list, + model_access_groups=model_access_groups, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + user_api_key_cache=user_api_key_cache, + ) + return all_models +async def _apply_user_models_filter( + all_models: List[str], + user_api_key_dict: "UserAPIKeyAuth", + proxy_model_list: List[str], + model_access_groups: Dict[str, List[str]], + prisma_client: Optional["PrismaClient"], + proxy_logging_obj: Optional["ProxyLogging"], + user_api_key_cache: Optional["DualCache"], +) -> List[str]: + """ + Intersect `all_models` with `LiteLLM_UserTable.models` (Personal + Models) for the user behind `user_api_key_dict`. + + Returns `all_models` unchanged when: + - the key has no associated user_id (master key, service accounts), + - prisma/cache are unavailable, + - the user object can't be loaded, + - the user object has no model restrictions, + - or the user explicitly opts in to `all-proxy-models`. + + Returns `[]` when the user has `no-default-models` (sentinel matches + `can_user_call_model` behavior at inference time). + """ + from litellm.proxy._types import SpecialModelNames + from litellm.proxy.auth.auth_checks import get_user_object + from litellm.proxy.auth.model_checks import ( + filter_models_by_user_access, + get_user_models, + ) + + if ( + not user_api_key_dict.user_id + or prisma_client is None + or user_api_key_cache is None + ): + return all_models + + try: + user_obj = await get_user_object( + user_id=user_api_key_dict.user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + # Mirror the swallow in user_api_key_auth.py — never break + # /v1/models if user lookup blips. No filter applied, which + # matches current behavior pre-fix. + verbose_proxy_logger.debug( + "_apply_user_models_filter: get_user_object failed, skipping " + "user-level filter. Exception: %s", + str(e), + ) + return all_models + + if user_obj is None or not user_obj.models: + return all_models + + if SpecialModelNames.no_default_models.value in user_obj.models: + return [] + + if SpecialModelNames.all_proxy_models.value in user_obj.models: + return all_models + + user_allowed = get_user_models( + user_models=list(user_obj.models), + proxy_model_list=proxy_model_list, + model_access_groups=model_access_groups, + ) + return filter_models_by_user_access( + models=all_models, + user_allowed_models=user_allowed, + ) + + def create_model_info_response( model_id: str, provider: str, diff --git a/tests/test_litellm/proxy/discovery_endpoints/test_v1_models_user_filter.py b/tests/test_litellm/proxy/discovery_endpoints/test_v1_models_user_filter.py new file mode 100644 index 00000000000..8166c25e662 --- /dev/null +++ b/tests/test_litellm/proxy/discovery_endpoints/test_v1_models_user_filter.py @@ -0,0 +1,240 @@ +""" +Route-level tests for /v1/models — verify that LiteLLM_UserTable.models +("Personal Models") is honored, closing the inconsistency between the +discovery endpoint and inference (BerriAI/litellm#26420). + +These tests exercise the real FastAPI route → real user_api_key_auth +dependency override → real model_list handler → real +get_available_models_for_user → real _apply_user_models_filter → +mocked get_user_object. No real DB, no live HTTP server. +""" + +import os +import sys +from unittest.mock import MagicMock + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm.proxy.proxy_server as ps +from litellm.caching.caching import DualCache +from litellm.proxy._types import ( + LiteLLM_UserTable, + LitellmUserRoles, + UserAPIKeyAuth, +) +from litellm.proxy.proxy_server import app + +# --------------------------------------------------------------------------- +# Test infrastructure +# --------------------------------------------------------------------------- + + +_PROXY_MODELS = [ + {"model_name": "gpt-4", "litellm_params": {"model": "openai/gpt-4"}}, + { + "model_name": "claude-3-opus", + "litellm_params": {"model": "anthropic/claude-3-opus"}, + }, + { + "model_name": "claude-3-haiku", + "litellm_params": {"model": "anthropic/claude-3-haiku"}, + }, + { + "model_name": "anthropic/claude-3-5-sonnet", + "litellm_params": {"model": "anthropic/claude-3-5-sonnet"}, + }, + { + "model_name": "anthropic/claude-3-7-sonnet", + "litellm_params": {"model": "anthropic/claude-3-7-sonnet"}, + }, +] + + +@pytest.fixture +def client(): + return TestClient(app) + + +@pytest.fixture +def configure_router(monkeypatch): + """Wire a real Router with a known model list onto proxy_server globals.""" + import litellm + + router = litellm.Router( + model_list=_PROXY_MODELS, + model_group_alias={}, + ) + monkeypatch.setattr(ps, "llm_router", router) + monkeypatch.setattr(ps, "llm_model_list", _PROXY_MODELS) + monkeypatch.setattr(ps, "user_model", None) + monkeypatch.setattr(ps, "general_settings", {}) + # _apply_user_models_filter only needs prisma_client to be truthy; + # the actual user fetch is mocked below. + monkeypatch.setattr(ps, "prisma_client", MagicMock()) + monkeypatch.setattr(ps, "user_api_key_cache", DualCache()) + return router + + +def _override_auth(user_id): + """Install an auth override that returns a UserAPIKeyAuth with the + given user_id (or None for master-key-style requests). The key itself + grants 'all-proxy-models' so key-level filtering is a no-op and we + can isolate the user-level filter under test. + """ + + def _auth(): + return UserAPIKeyAuth( + api_key="test-key", + user_id=user_id, + user_role=LitellmUserRoles.INTERNAL_USER, + models=[], # empty key.models → falls through to proxy list + team_id=None, + team_models=[], + ) + + app.dependency_overrides[ps.user_api_key_auth] = _auth + + +@pytest.fixture(autouse=True) +def _clear_auth_override(): + yield + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +def _patch_user(monkeypatch, models): + """Make get_user_object (as imported inside _apply_user_models_filter) + return a user with the given Personal Models list.""" + + async def _fake(*args, **kwargs): + return LiteLLM_UserTable( + user_id="u-test", + max_budget=None, + user_email=None, + models=models, + ) + + monkeypatch.setattr( + "litellm.proxy.auth.auth_checks.get_user_object", + _fake, + ) + + +def _model_ids(resp_json): + return [item["id"] for item in resp_json["data"]] + + +# --------------------------------------------------------------------------- +# Cases — these all hit GET /v1/models for real. +# --------------------------------------------------------------------------- + + +def test_v1_models_filters_by_user_personal_models( + client, configure_router, monkeypatch +): + """The headline bug: user.models = ['claude-3-opus'] must narrow the list.""" + _override_auth(user_id="u-test") + _patch_user(monkeypatch, models=["claude-3-opus"]) + + resp = client.get("/v1/models", headers={"Authorization": "Bearer sk-test"}) + assert resp.status_code == 200 + assert _model_ids(resp.json()) == ["claude-3-opus"] + + +def test_v1_models_no_filter_when_user_models_empty( + client, configure_router, monkeypatch +): + """user.models == [] → unrestricted, full proxy list returned.""" + _override_auth(user_id="u-test") + _patch_user(monkeypatch, models=[]) + + resp = client.get("/v1/models", headers={"Authorization": "Bearer sk-test"}) + assert resp.status_code == 200 + returned = set(_model_ids(resp.json())) + assert returned == {m["model_name"] for m in _PROXY_MODELS} + + +def test_v1_models_no_filter_when_user_id_is_none( + client, configure_router, monkeypatch +): + """Master key / service account (user_id=None) → no filter applied.""" + + async def _should_not_be_called(*args, **kwargs): + raise AssertionError("get_user_object must not be called when user_id is None") + + monkeypatch.setattr( + "litellm.proxy.auth.auth_checks.get_user_object", + _should_not_be_called, + ) + + _override_auth(user_id=None) + + resp = client.get("/v1/models", headers={"Authorization": "Bearer sk-test"}) + assert resp.status_code == 200 + assert len(_model_ids(resp.json())) == len(_PROXY_MODELS) + + +def test_v1_models_no_default_models_returns_empty( + client, configure_router, monkeypatch +): + """`no-default-models` sentinel → /v1/models returns [].""" + _override_auth(user_id="u-test") + _patch_user(monkeypatch, models=["no-default-models"]) + + resp = client.get("/v1/models", headers={"Authorization": "Bearer sk-test"}) + assert resp.status_code == 200 + assert resp.json()["data"] == [] + + +def test_v1_models_all_proxy_models_no_filter(client, configure_router, monkeypatch): + """`all-proxy-models` sentinel → user gets the full proxy list.""" + _override_auth(user_id="u-test") + _patch_user(monkeypatch, models=["all-proxy-models"]) + + resp = client.get("/v1/models", headers={"Authorization": "Bearer sk-test"}) + assert resp.status_code == 200 + returned = set(_model_ids(resp.json())) + assert returned == {m["model_name"] for m in _PROXY_MODELS} + + +def test_v1_models_user_wildcard(client, configure_router, monkeypatch): + """user.models contains 'anthropic/*' → only anthropic/* models returned.""" + _override_auth(user_id="u-test") + _patch_user(monkeypatch, models=["anthropic/*"]) + + resp = client.get("/v1/models", headers={"Authorization": "Bearer sk-test"}) + assert resp.status_code == 200 + returned = set(_model_ids(resp.json())) + assert returned == { + "anthropic/claude-3-5-sonnet", + "anthropic/claude-3-7-sonnet", + } + + +def test_v1_models_user_models_takes_precedence_over_permissive_key( + client, configure_router, monkeypatch +): + """ + Even if the key has 'all-proxy-models', user.models still narrows the + list. This is the parity claim: /v1/models matches can_user_call_model + at inference time. + """ + + def _auth_with_all_proxy(): + return UserAPIKeyAuth( + api_key="test-key", + user_id="u-test", + user_role=LitellmUserRoles.INTERNAL_USER, + models=["all-proxy-models"], + team_id=None, + team_models=[], + ) + + app.dependency_overrides[ps.user_api_key_auth] = _auth_with_all_proxy + _patch_user(monkeypatch, models=["claude-3-haiku"]) + + resp = client.get("/v1/models", headers={"Authorization": "Bearer sk-test"}) + assert resp.status_code == 200 + assert _model_ids(resp.json()) == ["claude-3-haiku"]