diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index fa979dcec39..0fc27c091d3 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index 01462a76730..806bfb7ba7d 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -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 # ──────────────────────────────────────────────