mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): apply user.models filter on /v1/models (BerriAI/litellm#26420)
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.
This commit is contained in:
parent
2b496bc7f7
commit
9fb55fe829
3 changed files with 395 additions and 0 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
Loading…
Add table
Reference in a new issue