mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
feat(proxy): auth_v2 identity enrichment so budget/limit hooks work for non-key logins
Virtual keys arrive fully populated via get_key_object; master/JWT/OAuth logins returned a thin identity, so the existing pre-call budget/limit hooks read None and enforced nothing for them. enrich_identity copies the user/team budget+limit fields 1:1 from the user/team rows into the identity's distinct user_*/team_* slots, filling only unset fields (never overriding an already-resolved value). Wired into the inference path for non-virtual-key logins, with get_user_object / get_team_object injected as loaders so the mapping is unit-tested without a DB. Additive by construction: these logins enforce nothing today, so it cannot regress existing behavior. The exact enforcement still needs a live rate-limit check before it is trusted; see INTEGRATION.md.
This commit is contained in:
parent
c0e06068e8
commit
205328bfb7
5 changed files with 191 additions and 14 deletions
|
|
@ -54,21 +54,19 @@ independent of v1/v2.
|
|||
JWT, and OAuth authenticators return a thin identity (no budget/limit fields), so
|
||||
the hooks read `None` and enforce nothing for those logins.
|
||||
|
||||
**Change.** Add an enrichment stage that runs after authentication: when the
|
||||
identity lacks budget/limit fields, load the user and team and populate them.
|
||||
**Built.** `enrichment.py` (`enrich_identity`) copies the user/team limit fields
|
||||
1:1 from the rows into the identity's distinct `user_*` / `team_*` slots, filling
|
||||
only unset fields so it never overrides an already-resolved value. It is wired
|
||||
into the inference path in `entry.py` for non-virtual-key logins
|
||||
(`_enrich_for_limits`), with the loaders (`get_user_object` auth_checks.py:1650,
|
||||
`get_team_object` auth_checks.py:1982) injected. Unit-tested in
|
||||
`test_enrichment.py`.
|
||||
|
||||
from litellm.proxy.auth.auth_checks import get_user_object, get_team_object
|
||||
# get_user_object: auth_checks.py:1650 get_team_object: auth_checks.py:1982
|
||||
|
||||
Populate exactly the fields the hooks read, sourced from the user/team rows.
|
||||
Dependency-inject the loaders so the stage is unit-testable without a DB.
|
||||
|
||||
**Why this is not done blind here.** The field mapping is subtle: a JWT principal
|
||||
has no key, so "the limit" is the user limit, the team limit, or their
|
||||
combination, and the hooks aggregate key/user/team in a specific way. Getting it
|
||||
wrong silently under- or over-enforces a customer's spend. It is additive today
|
||||
(these logins currently enforce nothing, so this cannot regress existing
|
||||
behavior), but the exact mapping must be confirmed on a running proxy.
|
||||
**What remains (live).** Only the verification below. The mapping is additive
|
||||
(these logins enforce nothing today, so it cannot regress existing behavior), but
|
||||
because it newly turns on enforcement for JWT/master/OAuth requests, the exact
|
||||
behavior must be confirmed against a running proxy before it is trusted — a wrong
|
||||
limit silently over- or under-enforces a customer's spend.
|
||||
|
||||
**Verify (live).**
|
||||
1. Create a user with `rpm_limit: 2`. Authenticate as that user via JWT.
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from .context import (
|
|||
try_get_auth_context,
|
||||
)
|
||||
from .end_user import resolve_end_user
|
||||
from .enrichment import enrich_identity
|
||||
from .entry import user_api_key_auth_v2
|
||||
from .telemetry import identity_span_attributes
|
||||
|
||||
|
|
@ -20,4 +21,5 @@ __all__ = [
|
|||
"attach_end_user",
|
||||
"resolve_end_user",
|
||||
"identity_span_attributes",
|
||||
"enrich_identity",
|
||||
]
|
||||
|
|
|
|||
61
litellm/proxy/auth/v2/enrichment.py
Normal file
61
litellm/proxy/auth/v2/enrichment.py
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
from typing import Any, Awaitable, Callable, Optional
|
||||
|
||||
# Loaders are injected so the mapping is unit-testable without a DB.
|
||||
UserLoader = Callable[[str], Awaitable[Optional[Any]]]
|
||||
TeamLoader = Callable[[str], Awaitable[Optional[Any]]]
|
||||
|
||||
# Source attr on the user/team row -> destination attr on the identity. The
|
||||
# destination user_*/team_* slots are distinct from key-level fields, so this
|
||||
# never overwrites a virtual key's own limits.
|
||||
_USER_FIELD_MAP = {
|
||||
"max_budget": "user_max_budget",
|
||||
"tpm_limit": "user_tpm_limit",
|
||||
"rpm_limit": "user_rpm_limit",
|
||||
"spend": "user_spend",
|
||||
}
|
||||
_TEAM_FIELD_MAP = {
|
||||
"max_budget": "team_max_budget",
|
||||
"tpm_limit": "team_tpm_limit",
|
||||
"rpm_limit": "team_rpm_limit",
|
||||
"spend": "team_spend",
|
||||
"models": "team_models",
|
||||
"blocked": "team_blocked",
|
||||
}
|
||||
|
||||
|
||||
def _copy_missing(identity: Any, source: Any, field_map: dict) -> None:
|
||||
for src_attr, dest_attr in field_map.items():
|
||||
value = getattr(source, src_attr, None)
|
||||
if value is not None and getattr(identity, dest_attr, None) is None:
|
||||
setattr(identity, dest_attr, value)
|
||||
|
||||
|
||||
async def enrich_identity(
|
||||
identity: Any,
|
||||
*,
|
||||
load_user: Optional[UserLoader] = None,
|
||||
load_team: Optional[TeamLoader] = None,
|
||||
) -> Any:
|
||||
"""Populate the identity's user/team budget+limit fields from the user/team rows.
|
||||
|
||||
Virtual keys arrive fully populated via ``get_key_object``; master/JWT/OAuth
|
||||
logins do not, so the existing pre-call budget/limit hooks read ``None`` and
|
||||
enforce nothing. This fills the gap, copying only fields that are unset (so it
|
||||
never overrides an already-resolved value) straight from the source rows.
|
||||
|
||||
Mechanically faithful and additive; the exact enforcement still needs a live
|
||||
rate-limit check before it is trusted (see INTEGRATION.md).
|
||||
"""
|
||||
user_id = getattr(identity, "user_id", None)
|
||||
if load_user is not None and user_id:
|
||||
user = await load_user(user_id)
|
||||
if user is not None:
|
||||
_copy_missing(identity, user, _USER_FIELD_MAP)
|
||||
|
||||
team_id = getattr(identity, "team_id", None)
|
||||
if load_team is not None and team_id:
|
||||
team = await load_team(team_id)
|
||||
if team is not None:
|
||||
_copy_missing(identity, team, _TEAM_FIELD_MAP)
|
||||
|
||||
return identity
|
||||
|
|
@ -7,6 +7,7 @@ from .authorizer import AuthorizationDenied, authorize
|
|||
from .context import AuthMethod, RequestAuthContext, set_auth_context
|
||||
from .end_user import resolve_end_user
|
||||
from .enforcer import CasbinEnforcer
|
||||
from .enrichment import enrich_identity
|
||||
from .policy_store import load_policy_snapshot
|
||||
from .principal import Principal, build_principal
|
||||
from .route_map import is_inference_route, match_route
|
||||
|
|
@ -38,6 +39,34 @@ async def _build_enforcer(principal: Principal, prisma_client: Any) -> CasbinEnf
|
|||
)
|
||||
|
||||
|
||||
async def _enrich_for_limits(identity: Any, ctx: AuthContext) -> None:
|
||||
"""Fill user/team budget+limit fields for non-key logins (master/JWT/OAuth) so
|
||||
the existing pre-call budget/limit hooks can enforce them. Virtual keys are
|
||||
already populated by get_key_object and skip this."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object, get_user_object
|
||||
|
||||
async def load_user(user_id: str) -> Any:
|
||||
return await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=ctx.prisma_client,
|
||||
user_api_key_cache=ctx.user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=ctx.parent_otel_span,
|
||||
proxy_logging_obj=ctx.proxy_logging_obj,
|
||||
)
|
||||
|
||||
async def load_team(team_id: str) -> Any:
|
||||
return await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=ctx.prisma_client,
|
||||
user_api_key_cache=ctx.user_api_key_cache,
|
||||
parent_otel_span=ctx.parent_otel_span,
|
||||
proxy_logging_obj=ctx.proxy_logging_obj,
|
||||
)
|
||||
|
||||
await enrich_identity(identity, load_user=load_user, load_team=load_team)
|
||||
|
||||
|
||||
async def _best_effort_identity(api_key: Optional[str], ctx: AuthContext) -> AuthResult:
|
||||
"""On loud-open routes, use the real identity if a usable key is present,
|
||||
otherwise fall back to an anonymous principal. Never fails the request."""
|
||||
|
|
@ -113,6 +142,8 @@ async def user_api_key_auth_v2(
|
|||
# key.models / access-group mechanism is intentionally not consulted.
|
||||
result = await authenticate(token, ctx)
|
||||
request_data = await _read_request_body(request=request)
|
||||
if result.method is not AuthMethod.VIRTUAL_KEY:
|
||||
await _enrich_for_limits(result.identity, ctx)
|
||||
principal, identity = _establish_context(request, result, route)
|
||||
requested_model = (
|
||||
request_data.get("model") if isinstance(request_data, dict) else None
|
||||
|
|
|
|||
85
tests/test_litellm/proxy/auth/v2/test_enrichment.py
Normal file
85
tests/test_litellm/proxy/auth/v2/test_enrichment.py
Normal file
|
|
@ -0,0 +1,85 @@
|
|||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.auth.v2.enrichment import enrich_identity
|
||||
|
||||
|
||||
def _identity(**overrides):
|
||||
base = dict(
|
||||
user_id="u1",
|
||||
team_id="t1",
|
||||
user_max_budget=None,
|
||||
user_tpm_limit=None,
|
||||
user_rpm_limit=None,
|
||||
user_spend=None,
|
||||
team_max_budget=None,
|
||||
team_tpm_limit=None,
|
||||
team_rpm_limit=None,
|
||||
team_spend=None,
|
||||
team_models=None,
|
||||
team_blocked=None,
|
||||
)
|
||||
base.update(overrides)
|
||||
return SimpleNamespace(**base)
|
||||
|
||||
|
||||
def _loader(obj):
|
||||
async def load(_id):
|
||||
return obj
|
||||
|
||||
return load
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_limits_are_copied_from_the_user_row():
|
||||
identity = _identity()
|
||||
user = SimpleNamespace(max_budget=12.5, tpm_limit=100, rpm_limit=10, spend=3.0)
|
||||
await enrich_identity(identity, load_user=_loader(user), load_team=_loader(None))
|
||||
assert identity.user_max_budget == 12.5
|
||||
assert identity.user_tpm_limit == 100
|
||||
assert identity.user_rpm_limit == 10
|
||||
assert identity.user_spend == 3.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_limits_and_models_are_copied_from_the_team_row():
|
||||
identity = _identity()
|
||||
team = SimpleNamespace(
|
||||
max_budget=50.0,
|
||||
tpm_limit=1000,
|
||||
rpm_limit=100,
|
||||
spend=9.0,
|
||||
models=["gpt-4o"],
|
||||
blocked=True,
|
||||
)
|
||||
await enrich_identity(identity, load_user=_loader(None), load_team=_loader(team))
|
||||
assert identity.team_max_budget == 50.0
|
||||
assert identity.team_rpm_limit == 100
|
||||
assert identity.team_models == ["gpt-4o"]
|
||||
assert identity.team_blocked is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_already_set_fields_are_not_overwritten():
|
||||
# A value resolved earlier (e.g. a key's own user limit) must win over the row.
|
||||
identity = _identity(user_rpm_limit=7)
|
||||
user = SimpleNamespace(max_budget=None, tpm_limit=None, rpm_limit=999, spend=None)
|
||||
await enrich_identity(identity, load_user=_loader(user), load_team=_loader(None))
|
||||
assert identity.user_rpm_limit == 7
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_ids_and_loaders_are_a_noop():
|
||||
identity = _identity(user_id=None, team_id=None)
|
||||
await enrich_identity(identity) # no loaders, no ids
|
||||
assert identity.user_max_budget is None
|
||||
assert identity.team_max_budget is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_loader_returning_none_leaves_identity_untouched():
|
||||
identity = _identity()
|
||||
await enrich_identity(identity, load_user=_loader(None), load_team=_loader(None))
|
||||
assert identity.user_max_budget is None
|
||||
assert identity.team_rpm_limit is None
|
||||
Loading…
Add table
Reference in a new issue