From 437f316a277c8a07e80a1732414778ca7fa27611 Mon Sep 17 00:00:00 2001 From: shivam Date: Thu, 21 May 2026 16:51:32 -0700 Subject: [PATCH] fix(jwt): make AUTO_REGISTER functional in prod; raise on missing winner MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two correctness fixes flagged by Greptile on the AUTO_REGISTER path: 1. generate_key_helper_fn was called without table_name="key". Without that, the helper falls into the user-upsert branch (table_name in (None, "user")) and tries to insert into LiteLLM_UserTable with user_id=None, which hits the NOT NULL @id constraint. AUTO_REGISTER would never have succeeded in production. Now passes table_name="key" explicitly, matching the /key/generate caller. 2. When the race loser refetches the winner's mapping and gets None (winner row concurrently deleted), the previous code returned None — and the caller in _resolve_jwt_to_virtual_key then fell through to less- restrictive team-based JWT auth, silently bypassing the configured AUTO_REGISTER policy. Now raises HTTP 503 so the caller retries against a stable state rather than getting unintended fallback access. Adds one test for the 503 winner-vanishes path. Co-Authored-By: Claude Opus 4.7 (1M context) --- litellm/proxy/auth/user_api_key_auth.py | 20 ++++++- .../proxy_unit_tests/test_jwt_key_mapping.py | 53 +++++++++++++++++++ 2 files changed, 71 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 0fc27c091d3..df950cb6a73 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -650,8 +650,14 @@ async def _auto_register_jwt_mapping( generate_key_helper_fn, ) + # ``table_name="key"`` is required: without it, generate_key_helper_fn + # falls into the user-upsert branch (`table_name is None or "user"`) and + # attempts to insert into LiteLLM_UserTable with user_id=None, which fails + # the NOT NULL @id constraint. Every successful key-creation caller (e.g. + # /key/generate) passes table_name="key" explicitly. key_data = await generate_key_helper_fn( request_type="key", + table_name="key", team_id=team_id, metadata={ "auto_registered": True, @@ -705,8 +711,18 @@ async def _auto_register_jwt_mapping( prisma_client=prisma_client, ) if token_hash is None: - # Should not happen, but guard against a delete racing our fetch. - return None + # The winner's mapping vanished between the unique-constraint + # conflict and our re-fetch (concurrent delete). Returning None + # here would silently fall through to team-based JWT auth — + # a less-restrictive path than the operator configured. Raise + # 503 so the caller retries against a stable state instead. + raise HTTPException( + status_code=503, + detail=( + "JWT Key Mapping: AUTO_REGISTER race resolution failed — " + "winner's mapping was concurrently removed. Retry the request." + ), + ) else: raise diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index 806bfb7ba7d..dc47f55a314 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -1069,6 +1069,59 @@ async def test_auto_register_race_conflict_tolerates_delete_failure(): prisma_client.db.litellm_verificationtoken.delete.assert_called_once() +@pytest.mark.asyncio +async def test_auto_register_raises_503_when_winner_mapping_vanishes(): + """ + Race edge case: this request loses the unique-constraint race, deletes its + orphan, then refetches the winner's mapping — but the winner's row was + concurrently deleted. Previously this returned None, silently falling + through to less-restrictive team-based JWT auth (bypassing the configured + AUTO_REGISTER policy). Must now raise HTTP 503 so the caller retries + rather than getting unintended fallback access. + """ + from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping + 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, + virtual_key_mapping_cache_ttl=300, + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_jwtkeymapping.create = AsyncMock( + side_effect=Exception("Unique constraint failed (P2002)") + ) + prisma_client.db.litellm_verificationtoken.delete = AsyncMock() + # Winner row no longer exists by the time we refetch + prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) + + user_api_key_cache = DualCache() + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + new_callable=AsyncMock, + return_value={"token": "sk-loser", "key": "sk-loser"}, + ), + pytest.raises(HTTPException) as exc_info, + ): + await _auto_register_jwt_mapping( + virtual_key_claim_field="sub", + claim_value="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:user-42", + ) + + assert exc_info.value.status_code == 503 + assert "concurrently removed" in exc_info.value.detail + + # ────────────────────────────────────────────── # Tests: AUTO_REGISTER inherits team_id from JWT # ──────────────────────────────────────────────