mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
7dcf18cfdf
commit
bf55895147
2 changed files with 197 additions and 39 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ──────────────────────────────────────────────
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue