fix(jwt): consistent AUTO_REGISTER on cached sentinel; clean up race orphans

Addresses Greptile review on PR #25570 cherry-pick.

1. Inconsistent AUTO_REGISTER when __NO_MAPPING__ sentinel is cached:
   The cached-sentinel branch silently returned None when prisma_client was
   None, while the fresh path raised HTTP 500 under the same config. Same
   request, different access-control outcome depending on cache state. Both
   paths now raise the same 500.

2. Orphaned virtual keys from race-condition losers:
   On unique-constraint conflict, generate_key_helper_fn had already persisted
   an unrestricted virtual key in LiteLLM_VerificationToken with the cleartext
   in request memory. Under sustained concurrency these accumulated
   indefinitely. The loser now deletes its orphan before falling back to the
   winner's mapping; failure to delete is logged but does not fail the request.

Also corrects a latent FK bug surfaced while fixing #2: the mapping row was
storing the plaintext key in LiteLLM_JWTKeyMapping.token, but that column FKs
to the hashed LiteLLM_VerificationToken.token — now hashed at the call site.

Tests:
- updated test_auto_register_creates_key_and_mapping to assert the hashed
  token is stored, not the plaintext
- updated test_auto_register_race_condition_unique_conflict to assert the
  orphan is deleted with the correct hashed token
- added test_auto_register_raises_500_when_sentinel_cached_and_no_db
- added test_auto_register_race_conflict_tolerates_delete_failure

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

View file

@ -620,7 +620,10 @@ async def _auto_register_jwt_mapping(
"jwt_claim_value": claim_value,
},
)
token_hash = key_data["token"]
# generate_key_helper_fn returns the plaintext key in "token"; the persisted
# row in LiteLLM_VerificationToken uses its hash, so hash here to get the FK
# value referenced by LiteLLM_JWTKeyMapping.token.
token_hash = hash_token(key_data["token"])
try:
await prisma_client.db.litellm_jwtkeymapping.create(
@ -635,15 +638,28 @@ async def _auto_register_jwt_mapping(
except Exception as e:
error_str = str(e).lower()
if "unique" in error_str or "p2002" in error_str:
# A concurrent request won the race — fetch the winning mapping and
# use its token. The key we just generated is orphaned but harmless;
# it will be excluded from spend tracking since nothing maps to it.
# A concurrent request won the race. The key generate_key_helper_fn
# just persisted to LiteLLM_VerificationToken is orphaned — nothing
# maps to it, but it's a fully valid unrestricted API key sitting in
# the DB and the cleartext is in memory on this request. Delete it
# so orphans don't accumulate under sustained concurrency.
verbose_proxy_logger.debug(
"JWT Key Mapping (auto_register): unique conflict on create — "
"fetching winner's mapping for %s='%s'.",
"deleting orphaned virtual key and fetching winner's mapping for %s='%s'.",
virtual_key_claim_field,
claim_value,
)
try:
await prisma_client.db.litellm_verificationtoken.delete(
where={"token": token_hash}
)
except Exception as delete_err:
# Don't fail the request if cleanup fails — the orphan is
# unmapped and inert. Log so an operator can prune it later.
verbose_proxy_logger.warning(
"JWT Key Mapping (auto_register): failed to delete orphaned key after race: %s",
delete_err,
)
token_hash = await get_jwt_key_mapping_object(
jwt_claim_name=virtual_key_claim_field,
jwt_claim_value=claim_value,
@ -710,9 +726,20 @@ async def _resolve_jwt_to_virtual_key(
status_code=403,
detail=f"JWT Key Mapping: No registered mapping for {virtual_key_claim_field}='{claim_value}'. Access denied.",
)
if behavior == UnregisteredJWTClientBehavior.AUTO_REGISTER and prisma_client is not None:
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.
# evict it and auto-register now that the policy has changed. 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.
if prisma_client is None:
raise HTTPException(
status_code=500,
detail=(
"JWT Key Mapping: AUTO_REGISTER requires a database connection. "
"Configure a database or change unregistered_jwt_client_behavior."
),
)
await user_api_key_cache.async_delete_cache(cache_key)
return await _auto_register_jwt_mapping(
virtual_key_claim_field=virtual_key_claim_field,

View file

@ -481,7 +481,9 @@ async def test_reject_behavior_raises_403_on_no_mapping():
user_api_key_cache = DualCache()
with patch("litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock):
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
):
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
@ -517,7 +519,9 @@ async def test_reject_behavior_caches_sentinel_after_db_miss():
user_api_key_cache = DualCache()
with patch("litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock):
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
):
# First call — DB miss, should raise 403 and write sentinel
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
@ -574,7 +578,9 @@ async def test_reject_behavior_raises_403_on_cached_no_mapping():
cache_key = "jwt_key_mapping:email:unknown@example.com"
await user_api_key_cache.async_set_cache(cache_key, "__NO_MAPPING__")
with patch("litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock):
with patch(
"litellm.proxy.auth.user_api_key_auth.get_key_object", new_callable=AsyncMock
):
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
@ -594,9 +600,10 @@ async def test_auto_register_creates_key_and_mapping():
"""
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.
UserAPIKeyAuth object. The mapping row stores the hashed token (FK to
LiteLLM_VerificationToken), not the plaintext key.
"""
from litellm.proxy._types import UnregisteredJWTClientBehavior
from litellm.proxy._types import UnregisteredJWTClientBehavior, hash_token
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
@ -611,15 +618,21 @@ async def test_auto_register_creates_key_and_mapping():
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock()
user_api_key_cache = DualCache()
mock_key_obj = UserAPIKeyAuth(token="hashed_auto_key", team_id=None)
plaintext_key = "sk-auto-key"
expected_hash = hash_token(plaintext_key)
mock_key_obj = UserAPIKeyAuth(token=expected_hash, 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"}
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": plaintext_key, "key": plaintext_key}
mock_get_key.return_value = mock_key_obj
result = await _resolve_jwt_to_virtual_key(
@ -632,15 +645,15 @@ async def test_auto_register_creates_key_and_mapping():
)
assert result == mock_key_obj
# Mapping row must have been created
# Mapping row must have been created with the hashed token (FK target)
prisma_client.db.litellm_jwtkeymapping.create.assert_called_once()
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"] == "hashed_auto_key"
# Cache must now hold the token hash
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 == "hashed_auto_key"
assert cached == expected_hash
@pytest.mark.asyncio
@ -672,12 +685,16 @@ async def test_auto_register_triggers_on_stale_no_mapping_sentinel():
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:
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
@ -699,11 +716,14 @@ async def test_auto_register_triggers_on_stale_no_mapping_sentinel():
async def test_auto_register_race_condition_unique_conflict():
"""
If two concurrent requests both call _auto_register_jwt_mapping and the
second hits a unique-constraint violation on create, it must fall back to
fetching the winner's mapping — no error surfaced to the caller.
second hits a unique-constraint violation on create, it must:
1) delete the orphaned virtual key it just created (so orphans don't
accumulate in LiteLLM_VerificationToken under sustained concurrency),
2) fall back to the winner's mapping,
3) not surface an error.
"""
from litellm.proxy.auth.user_api_key_auth import _auto_register_jwt_mapping
from litellm.proxy._types import UnregisteredJWTClientBehavior
from litellm.proxy._types import UnregisteredJWTClientBehavior, hash_token
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
@ -716,6 +736,7 @@ async def test_auto_register_race_condition_unique_conflict():
prisma_client.db.litellm_jwtkeymapping.create = AsyncMock(
side_effect=Exception("Unique constraint failed (P2002)")
)
prisma_client.db.litellm_verificationtoken.delete = AsyncMock()
# Simulate the winner's mapping already in DB after the conflict
winner_mapping = MagicMock()
winner_mapping.token = "winner_token_hash"
@ -725,14 +746,20 @@ async def test_auto_register_race_condition_unique_conflict():
)
user_api_key_cache = DualCache()
loser_plaintext = "sk-loser"
loser_hash = hash_token(loser_plaintext)
mock_key_obj = UserAPIKeyAuth(token="winner_token_hash", 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,
return_value={"token": "loser_token_hash", "key": "sk-loser"},
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,
return_value={"token": loser_plaintext, "key": loser_plaintext},
),
):
mock_get_key.return_value = mock_key_obj
@ -748,6 +775,10 @@ async def test_auto_register_race_condition_unique_conflict():
)
assert result == mock_key_obj
# The orphaned loser key must be deleted from LiteLLM_VerificationToken
prisma_client.db.litellm_verificationtoken.delete.assert_called_once_with(
where={"token": loser_hash}
)
# Cache should hold the winner's token, not the loser's
cached = await user_api_key_cache.async_get_cache("jwt_key_mapping:sub:user-42")
assert cached == "winner_token_hash"
@ -849,6 +880,106 @@ async def test_auto_register_raises_500_when_prisma_client_is_none():
assert "AUTO_REGISTER requires a database" in exc_info.value.detail
@pytest.mark.asyncio
async def test_auto_register_raises_500_when_sentinel_cached_and_no_db():
"""
AUTO_REGISTER + cached __NO_MAPPING__ sentinel + prisma_client is None must
raise HTTP 500, matching the fresh-path behavior. Previously this path
silently returned None and let the request fall through to team auth,
creating different access-control outcomes under identical configuration.
"""
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,
)
jwt_claims = {"sub": "user-42"}
user_api_key_cache = DualCache()
# Stale sentinel written under a prior fallback_team_mapping config
await user_api_key_cache.async_set_cache(
"jwt_key_mapping:sub:user-42", "__NO_MAPPING__"
)
with pytest.raises(HTTPException) as exc_info:
await _resolve_jwt_to_virtual_key(
jwt_claims=jwt_claims,
jwt_handler=jwt_handler,
prisma_client=None,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=None,
)
assert exc_info.value.status_code == 500
assert "AUTO_REGISTER requires a database" in exc_info.value.detail
@pytest.mark.asyncio
async def test_auto_register_race_conflict_tolerates_delete_failure():
"""
If deleting the orphaned virtual key after a race-condition conflict fails
(e.g. transient DB error), the request must still succeed by returning the
winner's mapping — the orphan is unmapped and inert.
"""
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(
side_effect=Exception("transient DB error")
)
winner_mapping = MagicMock()
winner_mapping.token = "winner_token_hash"
winner_mapping.is_active = True
prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(
return_value=winner_mapping
)
user_api_key_cache = DualCache()
mock_key_obj = UserAPIKeyAuth(token="winner_token_hash", 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,
return_value={"token": "sk-loser", "key": "sk-loser"},
),
):
mock_get_key.return_value = mock_key_obj
result = 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",
)
# Caller still receives the winner's mapping even when cleanup fails
assert result == mock_key_obj
prisma_client.db.litellm_verificationtoken.delete.assert_called_once()
# ──────────────────────────────────────────────
# Tests: backward-compat alias jwt_client_id_field
# ──────────────────────────────────────────────