mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
perf(proxy): skip redundant get_user_object on /v2/model/info?user_models_only=true
Closes the redundant-lookup path flagged by Greptile review on the PR: - non_admin_all_models already issues a litellm_usertable.find_unique for the calling user (to expand team membership). - apply_user_models_filter_to_deployments then called get_user_object for the same user_id immediately after. The second call hits the user_api_key_cache, but on cache miss it issued a second round-trip to the same row. Fix: non_admin_all_models now returns (deployments, user_models) so the caller in model_info_v2 can forward the already-loaded models list as user_models_override to both apply_user_models_filter_to_deployments and the underlying _apply_user_models_filter. The override short-circuits the get_user_object call entirely while preserving the existing SpecialModelNames / wildcard / access-group semantics. When override is None (every other code path — /v1/models, /v1/model/info, /v2/model/info without user_models_only, /v2/model/info?include_team_models without user_models_only), the behavior is identical to before. Tests (5 new parameterized cases in test_v2_model_info_user_filter.py): - get_user_object MUST NOT be called when override is passed - override=[] -> pass-through (no narrowing) - override=[no-default-models] -> [] - override=[all-proxy-models] -> pass-through - override=[anthropic/*] -> wildcard narrowing
This commit is contained in:
parent
42fe34cf4a
commit
b16d876a0f
3 changed files with 189 additions and 33 deletions
|
|
@ -11435,13 +11435,24 @@ async def non_admin_all_models(
|
|||
llm_router: Router,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: Optional[PrismaClient],
|
||||
):
|
||||
) -> Tuple[List[Dict], Optional[List[str]]]:
|
||||
"""
|
||||
Check if model is in db
|
||||
|
||||
Check if db model is 'created_by' == user_api_key_dict.user_id
|
||||
|
||||
Only return models that match
|
||||
|
||||
Returns
|
||||
-------
|
||||
(unique_models, user_models)
|
||||
unique_models: the de-duplicated deployment list as before.
|
||||
user_models: `LiteLLM_UserTable.models` (Personal Models) for
|
||||
the calling user if a row was loaded, else None.
|
||||
Lets the caller forward this to
|
||||
`apply_user_models_filter_to_deployments` as
|
||||
`user_models_override` to avoid a redundant
|
||||
`get_user_object` lookup downstream.
|
||||
"""
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -11456,6 +11467,7 @@ async def non_admin_all_models(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
user_models: Optional[List[str]] = None
|
||||
if user_api_key_dict.user_id:
|
||||
try:
|
||||
user_row = await UserRepository(prisma_client).table.find_unique(
|
||||
|
|
@ -11464,6 +11476,13 @@ async def non_admin_all_models(
|
|||
except Exception:
|
||||
raise HTTPException(status_code=400, detail={"error": "User not found"})
|
||||
|
||||
if user_row is not None:
|
||||
# Empty list "user has no model restriction configured" is
|
||||
# distinct from None "we never loaded a row" — keep the
|
||||
# distinction so the override path can short-circuit
|
||||
# correctly downstream.
|
||||
user_models = list(user_row.models or [])
|
||||
|
||||
# Get all models that are team models, when model team_id == user_row.teams
|
||||
all_models += _check_if_model_is_team_model(
|
||||
models=llm_router.get_model_list() or [],
|
||||
|
|
@ -11472,7 +11491,7 @@ async def non_admin_all_models(
|
|||
|
||||
# de-duplicate models. Only return unique model ids
|
||||
unique_models = _deduplicate_litellm_router_models(models=all_models)
|
||||
return unique_models
|
||||
return unique_models, user_models
|
||||
|
||||
|
||||
def _add_team_models_to_all_models(
|
||||
|
|
@ -12629,8 +12648,14 @@ async def model_info_v2(
|
|||
sort_by=sortBy,
|
||||
)
|
||||
|
||||
# When user_models_only=true, `non_admin_all_models` already
|
||||
# queries `litellm_usertable` to expand team-membership. Capture
|
||||
# the user's `models` list here so the filter step below can
|
||||
# forward it as `user_models_override` and skip a second
|
||||
# `get_user_object` call on cache miss.
|
||||
user_models_for_filter: Optional[List[str]] = None
|
||||
if user_models_only:
|
||||
all_models = await non_admin_all_models(
|
||||
all_models, user_models_for_filter = await non_admin_all_models(
|
||||
all_models=all_models,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
|
|
@ -12671,6 +12696,7 @@ async def model_info_v2(
|
|||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_models_override=user_models_for_filter,
|
||||
)
|
||||
|
||||
# Apply teamId filter if provided
|
||||
|
|
|
|||
|
|
@ -6462,6 +6462,7 @@ async def _apply_user_models_filter(
|
|||
prisma_client: Optional["PrismaClient"],
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
user_api_key_cache: Optional["DualCache"],
|
||||
user_models_override: Optional[List[str]] = None,
|
||||
) -> List[str]:
|
||||
"""
|
||||
Intersect `all_models` with `LiteLLM_UserTable.models` (Personal
|
||||
|
|
@ -6476,6 +6477,13 @@ async def _apply_user_models_filter(
|
|||
|
||||
Returns `[]` when the user has `no-default-models` (sentinel matches
|
||||
`can_user_call_model` behavior at inference time).
|
||||
|
||||
`user_models_override` lets callers that already loaded
|
||||
`LiteLLM_UserTable.models` (e.g. `non_admin_all_models` on
|
||||
`/v2/model/info?user_models_only=true`) pass the list in so we skip
|
||||
the redundant `get_user_object` DB hit on cache miss. Pass `[]`
|
||||
for "user has no model restrictions configured" and a populated
|
||||
list (including sentinels) when the user does.
|
||||
"""
|
||||
from litellm.proxy._types import SpecialModelNames
|
||||
from litellm.proxy.auth.auth_checks import get_user_object
|
||||
|
|
@ -6484,43 +6492,50 @@ async def _apply_user_models_filter(
|
|||
get_user_models,
|
||||
)
|
||||
|
||||
if (
|
||||
not user_api_key_dict.user_id
|
||||
or prisma_client is None
|
||||
or user_api_key_cache is None
|
||||
):
|
||||
if user_models_override is not None:
|
||||
user_models = list(user_models_override)
|
||||
else:
|
||||
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:
|
||||
return all_models
|
||||
user_models = list(user_obj.models or [])
|
||||
|
||||
if not user_models:
|
||||
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:
|
||||
if SpecialModelNames.no_default_models.value in user_models:
|
||||
return []
|
||||
|
||||
if SpecialModelNames.all_proxy_models.value in user_obj.models:
|
||||
if SpecialModelNames.all_proxy_models.value in user_models:
|
||||
return all_models
|
||||
|
||||
user_allowed = get_user_models(
|
||||
user_models=list(user_obj.models),
|
||||
user_models=user_models,
|
||||
proxy_model_list=proxy_model_list,
|
||||
model_access_groups=model_access_groups,
|
||||
)
|
||||
|
|
@ -6537,6 +6552,7 @@ async def apply_user_models_filter_to_deployments(
|
|||
prisma_client: Optional["PrismaClient"],
|
||||
proxy_logging_obj: Optional["ProxyLogging"],
|
||||
user_api_key_cache: Optional["DualCache"],
|
||||
user_models_override: Optional[List[str]] = None,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Apply the `LiteLLM_UserTable.models` (Personal Models) filter to a
|
||||
|
|
@ -6553,6 +6569,11 @@ async def apply_user_models_filter_to_deployments(
|
|||
|
||||
Order of `deployments` is preserved; duplicates with the same
|
||||
`model_name` are kept (multiple deployments can share a name).
|
||||
|
||||
`user_models_override` is forwarded to `_apply_user_models_filter`
|
||||
so callers that already loaded the row (e.g. `model_info_v2` on the
|
||||
`user_models_only=true` branch, where `non_admin_all_models` just
|
||||
queried `litellm_usertable`) skip the redundant lookup.
|
||||
"""
|
||||
if not deployments:
|
||||
return deployments
|
||||
|
|
@ -6575,6 +6596,7 @@ async def apply_user_models_filter_to_deployments(
|
|||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_models_override=user_models_override,
|
||||
)
|
||||
|
||||
allowed_set = set(allowed_model_names)
|
||||
|
|
|
|||
|
|
@ -346,3 +346,111 @@ def test_v2_model_info_include_team_models_no_default_models_returns_empty(
|
|||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["data"] == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct unit tests on apply_user_models_filter_to_deployments — verify the
|
||||
# `user_models_override` shortcut bypasses get_user_object entirely so the
|
||||
# `user_models_only=true` branch of /v2/model/info doesn't hit the DB twice
|
||||
# for the same user_id under cache-miss conditions.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_user_models_filter_override_narrows_and_skips_get_user_object(
|
||||
configure_router, monkeypatch
|
||||
):
|
||||
"""When the caller passes user_models_override, get_user_object MUST NOT
|
||||
be called — the override list is the single source of truth."""
|
||||
from litellm.proxy.utils import apply_user_models_filter_to_deployments
|
||||
|
||||
async def _must_not_call(*args, **kwargs):
|
||||
raise AssertionError(
|
||||
"get_user_object must not be called when user_models_override is set"
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.auth_checks.get_user_object",
|
||||
_must_not_call,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="u-test",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
models=[],
|
||||
team_id=None,
|
||||
team_models=[],
|
||||
)
|
||||
|
||||
result = await apply_user_models_filter_to_deployments(
|
||||
deployments=list(_PROXY_MODELS),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=configure_router,
|
||||
prisma_client=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_cache=DualCache(),
|
||||
user_models_override=["claude-3-opus"],
|
||||
)
|
||||
|
||||
assert sorted({d["model_name"] for d in result}) == ["claude-3-opus"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"override,expected",
|
||||
[
|
||||
# Empty list = "user has no model restriction" -> pass-through (no narrowing).
|
||||
([], None),
|
||||
# Sentinel -> empty.
|
||||
(["no-default-models"], []),
|
||||
# Sentinel -> pass-through (no narrowing).
|
||||
(["all-proxy-models"], None),
|
||||
# Wildcard.
|
||||
(
|
||||
["anthropic/*"],
|
||||
["anthropic/claude-3-5-sonnet", "anthropic/claude-3-7-sonnet"],
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_apply_user_models_filter_override_semantics(
|
||||
configure_router, monkeypatch, override, expected
|
||||
):
|
||||
"""The override path must apply the same SpecialModelNames + wildcard
|
||||
semantics as the get_user_object path, so swapping in the shortcut
|
||||
can never change observable behavior."""
|
||||
from litellm.proxy.utils import apply_user_models_filter_to_deployments
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.auth_checks.get_user_object",
|
||||
lambda *a, **kw: (_ for _ in ()).throw(
|
||||
AssertionError("override path must not consult get_user_object")
|
||||
),
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="u-test",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
models=[],
|
||||
team_id=None,
|
||||
team_models=[],
|
||||
)
|
||||
|
||||
result = await apply_user_models_filter_to_deployments(
|
||||
deployments=list(_PROXY_MODELS),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=configure_router,
|
||||
prisma_client=MagicMock(),
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_cache=DualCache(),
|
||||
user_models_override=override,
|
||||
)
|
||||
|
||||
if expected is None:
|
||||
# No narrowing — all five deployments returned, order-insensitive.
|
||||
assert sorted({d["model_name"] for d in result}) == sorted(
|
||||
{m["model_name"] for m in _PROXY_MODELS}
|
||||
)
|
||||
else:
|
||||
assert sorted({d["model_name"] for d in result}) == expected
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue