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:
shivam 2026-05-21 17:35:26 -07:00
parent 437f316a27
commit 76aaabc02d
No known key found for this signature in database
2 changed files with 189 additions and 252 deletions

View file

@ -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,

View file

@ -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"
# ──────────────────────────────────────────────