mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(jwt): defer AUTO_REGISTER until JWT policy is enforced by auth_builder
Closes the JWT policy bypass on the AUTO_REGISTER path flagged by veria-ai.
Before: when unregistered_jwt_client_behavior=auto_register and the JWT's
claim was unmapped, _resolve_jwt_to_virtual_key validated the JWT signature
and then immediately created a virtual key + mapping. JWTAuthManager.auth_builder
never ran for the first request (the new key short-circuited the team-auth
path), and every subsequent request hit the cached mapping — so custom_validate,
RBAC, scope_mappings, and user_allowed_email_domain were never enforced for
auto-registered clients.
After: _resolve_jwt_to_virtual_key returns a _PendingAutoRegister signal
instead of creating the key. The caller in _user_api_key_auth_builder runs
JWTAuthManager.auth_builder, then — only on a validated, policy-passing
result — calls _auto_register_jwt_mapping with the team_id / user_id from
that result. The created key inherits team + user limits from the validated
identity, and future cache hits load that already-policy-checked key.
Also drops the interim _resolve_inherited_team_id helper that pulled team_id
from raw JWT claims — same bypass risk; team_id now comes exclusively from
auth_builder.
Tests:
- Rewrote two existing tests to assert _resolve_jwt_to_virtual_key returns
_PendingAutoRegister (no key created yet) for both the fresh-DB-miss
and stale-sentinel branches
- Added a contract test that _auto_register_jwt_mapping stamps the
validated team_id/user_id onto generate_key_helper_fn
- Removed four stale team-binding tests that exercised the prior
raw-claim helper
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
437f316a27
commit
76aaabc02d
2 changed files with 189 additions and 252 deletions
|
|
@ -12,7 +12,7 @@ import fnmatch
|
|||
import re
|
||||
import secrets
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Iterator, List, Optional, Tuple, Union, cast
|
||||
from typing import Any, Iterator, List, NamedTuple, Optional, Tuple, Union, cast
|
||||
|
||||
import fastapi
|
||||
from fastapi import HTTPException, Request, WebSocket, status
|
||||
|
|
@ -587,38 +587,22 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints(
|
|||
return api_key
|
||||
|
||||
|
||||
def _resolve_inherited_team_id(
|
||||
jwt_handler: JWTHandler,
|
||||
jwt_claims: dict,
|
||||
) -> Optional[str]:
|
||||
class _PendingAutoRegister(NamedTuple):
|
||||
"""
|
||||
Resolve the team_id the JWT would have been bound to under standard team-
|
||||
based auth. AUTO_REGISTER stamps this on the new key so it inherits the
|
||||
team's budget/model/tpm/rpm/route limits — auto-registered clients then
|
||||
have the same constraints they'd have via the JWT team-auth path.
|
||||
Signal returned by ``_resolve_jwt_to_virtual_key`` when the JWT's claim is
|
||||
unmapped and ``unregistered_jwt_client_behavior`` is AUTO_REGISTER.
|
||||
|
||||
If ``team_id_jwt_field`` is configured but the JWT lacks the claim and no
|
||||
``team_id_default`` is set, raise 403: the operator has declared every JWT
|
||||
must carry a team, and registering a teamless (unbounded) key here would
|
||||
silently grant broader access than the team-auth path allows.
|
||||
The caller MUST run standard ``JWTAuthManager.auth_builder`` to apply RBAC,
|
||||
scope mappings, ``custom_validate``, and ``user_allowed_email_domain``
|
||||
policy BEFORE calling ``_auto_register_jwt_mapping`` with the validated
|
||||
``team_id`` / ``user_id`` from the auth_builder result. Auto-registering
|
||||
purely on a signature-valid JWT (the old behavior) bypassed every JWT
|
||||
policy beyond signature verification.
|
||||
"""
|
||||
if (
|
||||
jwt_handler.litellm_jwtauth.team_id_jwt_field is None
|
||||
and jwt_handler.litellm_jwtauth.team_id_default is None
|
||||
):
|
||||
return None
|
||||
|
||||
team_id = jwt_handler.get_team_id(token=jwt_claims, default_value=None)
|
||||
if team_id is None:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=(
|
||||
"JWT Key Mapping: AUTO_REGISTER requires a team_id on the JWT "
|
||||
"(via team_id_jwt_field) or a configured team_id_default — "
|
||||
"refusing to create an unbounded key. Access denied."
|
||||
),
|
||||
)
|
||||
return team_id
|
||||
claim_field: str
|
||||
claim_value: str
|
||||
cache_key: str
|
||||
|
||||
|
||||
async def _auto_register_jwt_mapping(
|
||||
|
|
@ -631,13 +615,15 @@ async def _auto_register_jwt_mapping(
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
cache_key: str,
|
||||
team_id: Optional[str] = None,
|
||||
user_id: Optional[str] = None,
|
||||
) -> Optional[UserAPIKeyAuth]:
|
||||
"""
|
||||
Auto-register: create a new virtual key + mapping for an unrecognised JWT
|
||||
claim value. The key is stamped with the JWT's resolved ``team_id`` (when
|
||||
available) so it inherits the team's budget/model/tpm/rpm/route limits via
|
||||
the standard auth path — auto-registered clients are then no more permissive
|
||||
than the team-based JWT auth path they would otherwise have taken.
|
||||
claim value. ``team_id`` and ``user_id`` must come from a successful
|
||||
``JWTAuthManager.auth_builder`` run — they encode the JWT identity AFTER
|
||||
RBAC/scope/custom_validate/email-domain policy has been enforced. The key
|
||||
is stamped with those values so the cached future-request path inherits
|
||||
the same team/user limits the auth_builder path would have applied.
|
||||
|
||||
Race safety: if two concurrent requests both reach here simultaneously (both
|
||||
saw no mapping in the DB), one will win the unique-constraint race on
|
||||
|
|
@ -659,6 +645,7 @@ async def _auto_register_jwt_mapping(
|
|||
request_type="key",
|
||||
table_name="key",
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
metadata={
|
||||
"auto_registered": True,
|
||||
"jwt_claim_field": virtual_key_claim_field,
|
||||
|
|
@ -754,7 +741,22 @@ async def _resolve_jwt_to_virtual_key(
|
|||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Optional[Span],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> Optional[UserAPIKeyAuth]:
|
||||
) -> Union[Optional[UserAPIKeyAuth], "_PendingAutoRegister"]:
|
||||
"""
|
||||
Returns:
|
||||
- ``UserAPIKeyAuth``: a resolved virtual key (cache hit or DB hit). The
|
||||
caller may use this directly; JWT policy has been enforced previously
|
||||
(at key-creation time or, for cached results, before caching).
|
||||
- ``_PendingAutoRegister``: claim is unmapped and behavior is AUTO_REGISTER.
|
||||
The caller MUST run ``JWTAuthManager.auth_builder`` to enforce JWT
|
||||
policy (RBAC, scope, custom_validate, email-domain), then invoke
|
||||
``_auto_register_jwt_mapping`` with the validated team_id/user_id.
|
||||
- ``None``: claim is unmapped and behavior is FALLBACK_TEAM_MAPPING.
|
||||
The caller falls through to standard team-based JWT auth (which itself
|
||||
enforces full JWT policy via auth_builder).
|
||||
- Raises HTTPException: REJECT policy hit, missing claim under
|
||||
REJECT/AUTO_REGISTER, or other policy violations.
|
||||
"""
|
||||
virtual_key_claim_field = jwt_handler.litellm_jwtauth.virtual_key_claim_field
|
||||
if virtual_key_claim_field is None:
|
||||
return None
|
||||
|
|
@ -800,7 +802,7 @@ async def _resolve_jwt_to_virtual_key(
|
|||
)
|
||||
if behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER:
|
||||
# Stale sentinel written under a prior fallback_team_mapping config —
|
||||
# evict it and auto-register now that the policy has changed. Raise
|
||||
# evict it and defer auto-register to after auth_builder runs. Raise
|
||||
# the same 500 as the fresh-path AUTO_REGISTER branch when there is
|
||||
# no DB, so behavior is consistent regardless of whether the cache
|
||||
# happens to hold the sentinel.
|
||||
|
|
@ -812,18 +814,11 @@ async def _resolve_jwt_to_virtual_key(
|
|||
"Configure a database or change unregistered_jwt_client_behavior."
|
||||
),
|
||||
)
|
||||
inherited_team_id = _resolve_inherited_team_id(jwt_handler, jwt_claims)
|
||||
await user_api_key_cache.async_delete_cache(cache_key)
|
||||
return await _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field=virtual_key_claim_field,
|
||||
return _PendingAutoRegister(
|
||||
claim_field=virtual_key_claim_field,
|
||||
claim_value=str(claim_value),
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
cache_key=cache_key,
|
||||
team_id=inherited_team_id,
|
||||
)
|
||||
return None
|
||||
elif cached_mapping is not None:
|
||||
|
|
@ -884,17 +879,14 @@ async def _resolve_jwt_to_virtual_key(
|
|||
"Configure a database or change unregistered_jwt_client_behavior."
|
||||
),
|
||||
)
|
||||
inherited_team_id = _resolve_inherited_team_id(jwt_handler, jwt_claims)
|
||||
return await _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field=virtual_key_claim_field,
|
||||
# Defer: caller runs JWTAuthManager.auth_builder to enforce RBAC, scope,
|
||||
# custom_validate, and email-domain policy, then auto-registers using
|
||||
# the validated identity. Auto-registering here on a signature-only
|
||||
# JWT would bypass every JWT policy beyond signature verification.
|
||||
return _PendingAutoRegister(
|
||||
claim_field=virtual_key_claim_field,
|
||||
claim_value=str(claim_value),
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
cache_key=cache_key,
|
||||
team_id=inherited_team_id,
|
||||
)
|
||||
|
||||
# FALLBACK_TEAM_MAPPING (default): cache the miss and return None so the
|
||||
|
|
@ -1093,6 +1085,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
# Try JWT-to-Virtual-Key mapping first to avoid
|
||||
# unnecessary DB queries in auth_builder
|
||||
do_standard_jwt_auth = True
|
||||
pending_auto_register: Optional[_PendingAutoRegister] = None
|
||||
if jwt_handler.litellm_jwtauth.virtual_key_claim_field is not None:
|
||||
# Decode JWT to get claims without running full auth_builder
|
||||
jwt_claims: Optional[dict]
|
||||
|
|
@ -1101,7 +1094,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
else:
|
||||
jwt_claims = await jwt_handler.auth_jwt(token=api_key)
|
||||
|
||||
valid_token = await _resolve_jwt_to_virtual_key(
|
||||
resolve_result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -1109,11 +1102,19 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if valid_token is not None:
|
||||
if isinstance(resolve_result, UserAPIKeyAuth):
|
||||
valid_token = resolve_result
|
||||
api_key = valid_token.token or ""
|
||||
valid_token.jwt_claims = jwt_claims
|
||||
do_standard_jwt_auth = False
|
||||
# Fall through to virtual key checks
|
||||
elif isinstance(resolve_result, _PendingAutoRegister):
|
||||
# Run full JWT policy (RBAC, scope, custom_validate,
|
||||
# email-domain) via auth_builder, then create the key
|
||||
# from the validated identity below.
|
||||
pending_auto_register = resolve_result
|
||||
# else: None → FALLBACK_TEAM_MAPPING, falls through to
|
||||
# standard JWT auth_builder below
|
||||
|
||||
if do_standard_jwt_auth:
|
||||
with tracer.trace("litellm.proxy.auth.jwt_auth_builder"):
|
||||
|
|
@ -1229,6 +1230,30 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
else None
|
||||
)
|
||||
|
||||
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
|
||||
# JWT policy (RBAC, scope, custom_validate, email-domain)
|
||||
# has now been enforced by auth_builder above. Create the
|
||||
# mapping + virtual key from the *validated* identity, then
|
||||
# replace valid_token with the new key so downstream checks
|
||||
# use the key-scoped path.
|
||||
if pending_auto_register is not None and prisma_client is not None:
|
||||
auto_registered = await _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field=pending_auto_register.claim_field,
|
||||
claim_value=pending_auto_register.claim_value,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
cache_key=pending_auto_register.cache_key,
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
)
|
||||
if auto_registered is not None:
|
||||
auto_registered.jwt_claims = jwt_claims
|
||||
valid_token = auto_registered
|
||||
api_key = valid_token.token or ""
|
||||
|
||||
# Check if model has zero cost - if so, skip all budget checks
|
||||
model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
|
|
|
|||
|
|
@ -596,14 +596,17 @@ async def test_reject_behavior_raises_403_on_cached_no_mapping():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_creates_key_and_mapping():
|
||||
async def test_auto_register_returns_pending_signal_without_creating_key():
|
||||
"""
|
||||
When unregistered_jwt_client_behavior='auto_register' and no mapping exists,
|
||||
_resolve_jwt_to_virtual_key must create a key + mapping row and return a
|
||||
UserAPIKeyAuth object. The mapping row stores the hashed token (FK to
|
||||
LiteLLM_VerificationToken), not the plaintext key.
|
||||
Security: when unregistered_jwt_client_behavior='auto_register' and no
|
||||
mapping exists, _resolve_jwt_to_virtual_key must NOT create the key yet.
|
||||
It returns a _PendingAutoRegister signal so the caller can run
|
||||
JWTAuthManager.auth_builder (enforcing RBAC, scope mappings,
|
||||
custom_validate, user_allowed_email_domain) FIRST. Creating the key here
|
||||
would bypass every JWT policy beyond signature verification.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior, hash_token
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
from litellm.proxy.auth.user_api_key_auth import _PendingAutoRegister
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
|
|
@ -617,10 +620,55 @@ async def test_auto_register_creates_key_and_mapping():
|
|||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
user_api_key_cache = DualCache()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key:
|
||||
result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert isinstance(result, _PendingAutoRegister)
|
||||
assert result.claim_field == "sub"
|
||||
assert result.claim_value == "new-user-42"
|
||||
assert result.cache_key == "jwt_key_mapping:sub:new-user-42"
|
||||
# CRITICAL: no key was created — that must wait until after auth_builder
|
||||
mock_gen_key.assert_not_called()
|
||||
prisma_client.db.litellm_jwtkeymapping.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_creates_key_and_mapping_when_helper_invoked():
|
||||
"""
|
||||
When the caller invokes _auto_register_jwt_mapping directly (after
|
||||
auth_builder validation), the helper creates the key + mapping row and
|
||||
returns a UserAPIKeyAuth. The mapping row stores the hashed token (FK to
|
||||
LiteLLM_VerificationToken), not the plaintext key.
|
||||
"""
|
||||
from litellm.proxy._types import hash_token
|
||||
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
virtual_key_mapping_cache_ttl=300,
|
||||
)
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
user_api_key_cache = DualCache()
|
||||
plaintext_key = "sk-auto-key"
|
||||
expected_hash = hash_token(plaintext_key)
|
||||
mock_key_obj = UserAPIKeyAuth(token=expected_hash, team_id=None)
|
||||
mock_key_obj = UserAPIKeyAuth(token=expected_hash, team_id="validated-team")
|
||||
|
||||
with (
|
||||
patch(
|
||||
|
|
@ -635,35 +683,45 @@ async def test_auto_register_creates_key_and_mapping():
|
|||
mock_gen_key.return_value = {"token": plaintext_key, "key": plaintext_key}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
result = await _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field="sub",
|
||||
claim_value="new-user-42",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
cache_key="jwt_key_mapping:sub:new-user-42",
|
||||
team_id="validated-team",
|
||||
user_id="validated-user",
|
||||
)
|
||||
|
||||
assert result == mock_key_obj
|
||||
# Mapping row must have been created with the hashed token (FK target)
|
||||
prisma_client.db.litellm_jwtkeymapping.create.assert_called_once()
|
||||
# generate_key_helper_fn was passed table_name="key" (not user-upsert path)
|
||||
# and the validated team_id + user_id from auth_builder
|
||||
assert mock_gen_key.call_args.kwargs["table_name"] == "key"
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] == "validated-team"
|
||||
assert mock_gen_key.call_args.kwargs["user_id"] == "validated-user"
|
||||
# Mapping row was created with the hashed token (FK target)
|
||||
call_data = prisma_client.db.litellm_jwtkeymapping.create.call_args[1]["data"]
|
||||
assert call_data["jwt_claim_name"] == "sub"
|
||||
assert call_data["jwt_claim_value"] == "new-user-42"
|
||||
assert call_data["token"] == expected_hash
|
||||
# Cache must hold the hashed token
|
||||
cached = await user_api_key_cache.async_get_cache("jwt_key_mapping:sub:new-user-42")
|
||||
assert cached == expected_hash
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_triggers_on_stale_no_mapping_sentinel():
|
||||
async def test_auto_register_returns_pending_signal_on_stale_no_mapping_sentinel():
|
||||
"""
|
||||
If the cache holds a stale __NO_MAPPING__ sentinel (written under a prior
|
||||
fallback_team_mapping config) and behavior is now AUTO_REGISTER, the sentinel
|
||||
must be evicted and auto-registration must run — not silently return None.
|
||||
fallback_team_mapping config) and behavior is now AUTO_REGISTER, the
|
||||
resolver must evict the sentinel and return _PendingAutoRegister (so the
|
||||
caller can run auth_builder before creating the key) — not silently return
|
||||
None and not create the key on the spot.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
from litellm.proxy.auth.user_api_key_auth import _PendingAutoRegister
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
|
|
@ -678,26 +736,14 @@ async def test_auto_register_triggers_on_stale_no_mapping_sentinel():
|
|||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
user_api_key_cache = DualCache()
|
||||
# Seed the stale sentinel
|
||||
await user_api_key_cache.async_set_cache(
|
||||
"jwt_key_mapping:email:alice@corp.com", "__NO_MAPPING__"
|
||||
)
|
||||
|
||||
mock_key_obj = UserAPIKeyAuth(token="hashed_auto_key", team_id=None)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
):
|
||||
mock_gen_key.return_value = {"token": "hashed_auto_key", "key": "sk-auto-key"}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key:
|
||||
result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
|
|
@ -707,9 +753,15 @@ async def test_auto_register_triggers_on_stale_no_mapping_sentinel():
|
|||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
# Must have auto-registered, not returned None
|
||||
assert result == mock_key_obj
|
||||
prisma_client.db.litellm_jwtkeymapping.create.assert_called_once()
|
||||
assert isinstance(result, _PendingAutoRegister)
|
||||
# Stale sentinel must be evicted so the deferred auto-register actually
|
||||
# runs after auth_builder validates the JWT
|
||||
cached_after = await user_api_key_cache.async_get_cache(
|
||||
"jwt_key_mapping:email:alice@corp.com"
|
||||
)
|
||||
assert cached_after is None
|
||||
mock_gen_key.assert_not_called()
|
||||
prisma_client.db.litellm_jwtkeymapping.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1123,33 +1175,33 @@ async def test_auto_register_raises_503_when_winner_mapping_vanishes():
|
|||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# Tests: AUTO_REGISTER inherits team_id from JWT
|
||||
# Tests: AUTO_REGISTER stamps validated identity from auth_builder
|
||||
# ──────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_stamps_team_id_from_jwt_claim():
|
||||
async def test_auto_register_helper_stamps_validated_team_and_user():
|
||||
"""
|
||||
When team_id_jwt_field is configured and the JWT carries the claim,
|
||||
AUTO_REGISTER must stamp the resolved team_id on the new key so it inherits
|
||||
the team's limits via standard auth (budget, models, tpm/rpm, allowed_routes).
|
||||
The deferred-auto-register contract: _auto_register_jwt_mapping is called
|
||||
with team_id and user_id from JWTAuthManager.auth_builder's *validated*
|
||||
result (after RBAC, scope mappings, custom_validate, email-domain policy).
|
||||
These must be passed to generate_key_helper_fn so the created key carries
|
||||
them — the cached future-request path then inherits the same team/user
|
||||
limits the auth_builder path would have applied.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
team_id_jwt_field="team_id",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
virtual_key_mapping_cache_ttl=300,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user", "team_id": "team-engineering"}
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
user_api_key_cache = DualCache()
|
||||
mock_key_obj = UserAPIKeyAuth(token="hashed", team_id="team-engineering")
|
||||
mock_key_obj = UserAPIKeyAuth(
|
||||
token="hashed", team_id="validated-team", user_id="validated-user"
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
|
|
@ -1164,163 +1216,23 @@ async def test_auto_register_stamps_team_id_from_jwt_claim():
|
|||
mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
result = await _auto_register_jwt_mapping(
|
||||
virtual_key_claim_field="sub",
|
||||
claim_value="new-user",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
cache_key="jwt_key_mapping:sub:new-user",
|
||||
team_id="validated-team",
|
||||
user_id="validated-user",
|
||||
)
|
||||
|
||||
assert result == mock_key_obj
|
||||
# generate_key_helper_fn must have been called with team_id from the JWT,
|
||||
# so the key inherits team-level limits via the standard auth path.
|
||||
mock_gen_key.assert_called_once()
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] == "team-engineering"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_rejects_when_team_id_required_but_missing_from_jwt():
|
||||
"""
|
||||
Security: when team_id_jwt_field is configured (operator wants every JWT
|
||||
bound to a team), AUTO_REGISTER must refuse a JWT that lacks the team
|
||||
claim. Otherwise the auto-created key would be teamless and inherit no
|
||||
budget/model/rate limits — broader access than the team-auth path allows.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
team_id_jwt_field="team_id",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user"} # no "team_id"
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
pytest.raises(HTTPException) as exc_info,
|
||||
):
|
||||
await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "team_id" in exc_info.value.detail
|
||||
# No key was created — we refused before generate_key_helper_fn ran
|
||||
mock_gen_key.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_uses_team_id_default_when_jwt_lacks_team_claim():
|
||||
"""
|
||||
When team_id_jwt_field is configured but the JWT lacks the claim and
|
||||
team_id_default is set, the default team_id is stamped on the new key —
|
||||
the explicit operator-chosen fallback team bounds the auto-registered key.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
team_id_jwt_field="team_id",
|
||||
team_id_default="default-team",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user"} # no "team_id"
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
mock_key_obj = UserAPIKeyAuth(token="hashed", team_id="default-team")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
):
|
||||
mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
result = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
assert result == mock_key_obj
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] == "default-team"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_register_no_team_id_when_team_field_not_configured():
|
||||
"""
|
||||
When the operator has NOT configured team_id_jwt_field or team_id_default,
|
||||
they have explicitly opted out of team-based bounding. AUTO_REGISTER then
|
||||
creates a teamless key (preserving prior behavior) — there is no team to
|
||||
inherit from. The key being unbounded matches what the team-auth path
|
||||
would have produced in the same config.
|
||||
"""
|
||||
from litellm.proxy._types import UnregisteredJWTClientBehavior
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
virtual_key_claim_field="sub",
|
||||
unregistered_jwt_client_behavior=UnregisteredJWTClientBehavior.AUTO_REGISTER,
|
||||
)
|
||||
jwt_claims = {"sub": "new-user", "team_id": "ignored-because-not-configured"}
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None)
|
||||
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
|
||||
|
||||
mock_key_obj = UserAPIKeyAuth(token="hashed", team_id=None)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_gen_key,
|
||||
):
|
||||
mock_gen_key.return_value = {"token": "sk-newkey", "key": "sk-newkey"}
|
||||
mock_get_key.return_value = mock_key_obj
|
||||
|
||||
await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=DualCache(),
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
|
||||
# team_id is None because the operator opted out of team-based config
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] is None
|
||||
# Validated identity flowed to the key (key inherits team + user limits)
|
||||
assert mock_gen_key.call_args.kwargs["team_id"] == "validated-team"
|
||||
assert mock_gen_key.call_args.kwargs["user_id"] == "validated-user"
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue