fix(jwt): AUTO_REGISTER inherits team_id so keys are bounded by team limits

Auto-registered virtual keys were created with no team, model, route, rate, or
budget constraints — broader access than the standard team-based JWT auth path
the same client would have taken. Under AUTO_REGISTER, resolve the team_id
from the JWT (via the operator-configured team_id_jwt_field / team_id_default)
and stamp it on the new key. Downstream auth then applies the team's
budget/models/tpm/rpm/allowed_routes via the existing virtual-key flow.

Policy when team_id_jwt_field is configured:
  - JWT carries team claim → stamp resolved team_id
  - JWT lacks claim + team_id_default set → stamp default
  - JWT lacks claim + no default → 403 (refuse to create an unbounded key)

When neither team_id_jwt_field nor team_id_default is configured, the
operator has explicitly opted out of team-based limits — the auto-created
key has no team_id (matches what team-auth would do in the same config).

Adds 4 tests covering each branch.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
shivam 2026-05-21 16:26:06 -07:00
parent 2c23811880
commit 2c42c5bd4b
No known key found for this signature in database
2 changed files with 250 additions and 6 deletions

View file

@ -587,6 +587,40 @@ 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]:
"""
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.
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.
"""
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
async def _auto_register_jwt_mapping(
virtual_key_claim_field: str,
claim_value: str,
@ -596,15 +630,19 @@ async def _auto_register_jwt_mapping(
parent_otel_span: Optional[Span],
proxy_logging_obj: ProxyLogging,
cache_key: str,
team_id: Optional[str] = None,
) -> Optional[UserAPIKeyAuth]:
"""
Auto-register: create a new virtual key + mapping for an unrecognised JWT claim value.
The new key carries no model/budget restrictions; admins can tighten it later.
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.
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
litellm_jwtkeymapping. The loser catches the conflict, fetches the winner's
mapping, and proceeds — no orphaned keys and no error surfaced to the caller.
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
litellm_jwtkeymapping. The loser catches the conflict, deletes its orphaned
key, fetches the winner's mapping, and proceeds — no error surfaced.
"""
# Inline import required: key_management_endpoints imports user_api_key_auth
# (line 51) so a module-level import here would create a circular dependency.
@ -614,6 +652,7 @@ async def _auto_register_jwt_mapping(
key_data = await generate_key_helper_fn(
request_type="key",
team_id=team_id,
metadata={
"auto_registered": True,
"jwt_claim_field": virtual_key_claim_field,
@ -757,6 +796,7 @@ 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,
@ -767,6 +807,7 @@ async def _resolve_jwt_to_virtual_key(
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:
@ -827,6 +868,7 @@ 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,
claim_value=str(claim_value),
@ -836,6 +878,7 @@ async def _resolve_jwt_to_virtual_key(
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

View file

@ -1069,6 +1069,207 @@ async def test_auto_register_race_conflict_tolerates_delete_failure():
prisma_client.db.litellm_verificationtoken.delete.assert_called_once()
# ──────────────────────────────────────────────
# Tests: AUTO_REGISTER inherits team_id from JWT
# ──────────────────────────────────────────────
@pytest.mark.asyncio
async def test_auto_register_stamps_team_id_from_jwt_claim():
"""
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).
"""
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", "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")
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=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
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
# ──────────────────────────────────────────────
# Tests: backward-compat alias jwt_client_id_field
# ──────────────────────────────────────────────