diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 52943737eed..6b1d845ec95 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1620,6 +1620,34 @@ async def _get_fuzzy_user_object( return response +async def _backfill_null_user_email( + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + user_row: LiteLLM_UserTable, + user_email: str | None, +) -> LiteLLM_UserTable: + if user_email is None or user_row.user_email is not None or prisma_client is None: + return user_row + + user_repo = UserRepository(prisma_client) + await user_repo.backfill_null_user_email( + user_id=user_row.user_id, + user_email=user_email, + ) + db_row = await user_repo.find_by_id(user_row.user_id) + if db_row is None: + return user_row + email_update = {"user_email": db_row.user_email} # mutable-ok: model_copy update payload is dict-shaped + updated_row = user_row.model_copy(update=email_update) + await user_api_key_cache.async_set_cache( + key=user_row.user_id, + value=updated_row, + model_type=LiteLLM_UserTable, + ttl=get_management_object_ttl(user_api_key_cache), + ) + return updated_row + + @log_db_metrics async def get_user_object( user_id: str | None, @@ -1648,7 +1676,12 @@ async def get_user_object( model_type=LiteLLM_UserTable, ) if cached_user_obj is not None: - return cached_user_obj + return await _backfill_null_user_email( + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_row=cached_user_obj, + user_email=user_email, + ) # else, check db if prisma_client is None: raise Exception("No db connected") @@ -1732,6 +1765,12 @@ async def get_user_object( response.organization_memberships = _dumped_memberships _response = LiteLLM_UserTable.model_validate(dict(response)) + _response = await _backfill_null_user_email( + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_row=_response, + user_email=user_email, + ) response_dict = _response.model_dump() # save the user object to cache diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 2905eb86c0f..286837c8909 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1232,6 +1232,26 @@ async def _user_api_key_auth_builder( valid_token.jwt_claims = jwt_claims do_standard_jwt_auth = False # Fall through to virtual key checks + if valid_token.user_id is not None and valid_token.user_email is None: + mapped_claims = jwt_claims or {} # mutable-ok: empty-dict fallback for the None-claims case + mapped_user_email = jwt_handler.get_user_email(token=mapped_claims, default_value=None) + mapped_jwt_user_id = jwt_handler.get_user_id(token=mapped_claims, default_value=None) + if mapped_user_email is not None and mapped_jwt_user_id == valid_token.user_id: + try: + mapped_user_obj = await get_user_object( + user_id=valid_token.user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + user_email=mapped_user_email, + ) + except Exception as e: + verbose_proxy_logger.debug(f"JWT mapped-key user_email backfill skipped: {e}") + else: + if mapped_user_obj is not None: + valid_token.user_email = mapped_user_obj.user_email elif isinstance(resolve_result, _PendingAutoRegister): # Run full JWT policy (RBAC, scope, custom_validate, # email-domain) via auth_builder, then create the key diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index 5eb326bda18..2b567e8b52a 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -195,6 +195,17 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): return await self.update(user_id, data, id_field="user_id") + async def backfill_null_user_email(self, user_id: str, user_email: str) -> int: + """Set user_email only when the stored value is null, atomically at the database. + + Returns the number of rows updated: 0 means another writer already set an email. + """ + updated_count: int = await self.table.update_many( + where={"user_id": user_id, "user_email": None}, # mutable-ok: Prisma query filters are dict-shaped + data={"user_email": user_email}, # mutable-ok: Prisma update payloads are dict-shaped + ) + return updated_count + async def delete_user(self, user_id: str) -> LiteLLM_UserTable | None: """Delete a user.""" return await self.delete(user_id, id_field="user_id") diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 5f3b0f36b95..f5aa695cb78 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -883,6 +883,229 @@ async def test_get_user_object_upsert_includes_user_email(): assert creation_args["user_id"] == "new_test_user" +@pytest.mark.asyncio +async def test_get_user_object_backfills_null_email_from_cache_hit(): + """ + Regression (LIT-4710): an existing user row with a null user_email must be + backfilled from the JWT-provided email even when served from cache, so the + JWT-to-virtual-key path (which resolves straight to the cached user) stops + logging user_api_key_user_email=null forever. Before the fix the cached row + was returned unchanged and the DB was never updated. + """ + cache = UserApiKeyCache() + existing = LiteLLM_UserTable( + user_id="jwt-user-1", user_email=None, user_role="internal_user" + ) + await cache.async_set_cache( + key="jwt-user-1", value=existing, model_type=LiteLLM_UserTable + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable( + user_id="jwt-user-1", + user_email="jwt-user-1@example.com", + user_role="internal_user", + ) + ) + + result = await get_user_object( + user_id="jwt-user-1", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + user_id_upsert=False, + proxy_logging_obj=None, + user_email="jwt-user-1@example.com", + ) + + assert result is not None + assert result.user_email == "jwt-user-1@example.com" + + mock_prisma_client.db.litellm_usertable.update_many.assert_called_once() + update_kwargs = mock_prisma_client.db.litellm_usertable.update_many.call_args.kwargs + assert update_kwargs["where"] == {"user_id": "jwt-user-1", "user_email": None} + assert update_kwargs["data"]["user_email"] == "jwt-user-1@example.com" + + refreshed = await cache.async_get_cache( + key="jwt-user-1", model_type=LiteLLM_UserTable + ) + assert refreshed is not None + assert refreshed.user_email == "jwt-user-1@example.com" + + +@pytest.mark.asyncio +async def test_get_user_object_backfills_null_email_from_db_read(): + """ + Regression (LIT-4710): a user row read from the DB with a null user_email is + backfilled from the JWT-provided email before it is cached and returned. + """ + cache = UserApiKeyCache() + db_row = LiteLLM_UserTable( + user_id="jwt-user-3", user_email=None, user_role="internal_user" + ) + backfilled_row = LiteLLM_UserTable( + user_id="jwt-user-3", + user_email="jwt-user-3@example.com", + user_role="internal_user", + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + side_effect=[db_row, backfilled_row] + ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1) + + with patch( + "litellm.proxy.auth.auth_checks._should_check_db", return_value=True + ): + result = await get_user_object( + user_id="jwt-user-3", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + user_id_upsert=False, + proxy_logging_obj=None, + user_email="jwt-user-3@example.com", + ) + + assert result is not None + assert result.user_email == "jwt-user-3@example.com" + mock_prisma_client.db.litellm_usertable.update_many.assert_called_once() + + refreshed = await cache.async_get_cache( + key="jwt-user-3", model_type=LiteLLM_UserTable + ) + assert refreshed is not None + assert refreshed.user_email == "jwt-user-3@example.com" + + +@pytest.mark.asyncio +async def test_get_user_object_does_not_overwrite_existing_email(): + """ + LIT-4710 guardrail: backfill is scoped to null-to-value. An existing non-null + user_email (e.g. one an operator set intentionally) must never be overwritten + by the JWT-provided email. + """ + cache = UserApiKeyCache() + existing = LiteLLM_UserTable( + user_id="jwt-user-2", + user_email="operator-set@example.com", + user_role="internal_user", + ) + await cache.async_set_cache( + key="jwt-user-2", value=existing, model_type=LiteLLM_UserTable + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0) + + result = await get_user_object( + user_id="jwt-user-2", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + user_id_upsert=False, + proxy_logging_obj=None, + user_email="different@example.com", + ) + + assert result is not None + assert result.user_email == "operator-set@example.com" + mock_prisma_client.db.litellm_usertable.update_many.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_user_object_backfill_race_prefers_db_email(): + """ + LIT-4710 race guard: when the null-guarded update matches 0 rows because a + concurrent writer already backfilled an email, the cache must be refreshed + with the value the DB accepted, not this request's proposed email. + """ + cache = UserApiKeyCache() + existing = LiteLLM_UserTable( + user_id="jwt-user-4", user_email=None, user_role="internal_user" + ) + await cache.async_set_cache( + key="jwt-user-4", value=existing, model_type=LiteLLM_UserTable + ) + + winner_row = LiteLLM_UserTable( + user_id="jwt-user-4", + user_email="winner@example.com", + user_role="internal_user", + ) + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=winner_row + ) + + result = await get_user_object( + user_id="jwt-user-4", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + user_id_upsert=False, + proxy_logging_obj=None, + user_email="loser@example.com", + ) + + assert result is not None + assert result.user_email == "winner@example.com" + + refreshed = await cache.async_get_cache( + key="jwt-user-4", model_type=LiteLLM_UserTable + ) + assert refreshed is not None + assert refreshed.user_email == "winner@example.com" + + +@pytest.mark.asyncio +async def test_get_user_object_backfill_caches_persisted_email_not_proposed(): + """ + LIT-4710 cache-coherence: even when the null-guarded update succeeds, the + cache must be refreshed from the row the DB actually holds, not this + request's proposed email. A concurrent ordinary user update (not null + guarded) can change the email in the window before the cache write, so + optimistically caching the proposed email would serve a stale value. + """ + cache = UserApiKeyCache() + existing = LiteLLM_UserTable( + user_id="jwt-user-5", user_email=None, user_role="internal_user" + ) + await cache.async_set_cache( + key="jwt-user-5", value=existing, model_type=LiteLLM_UserTable + ) + + persisted_row = LiteLLM_UserTable( + user_id="jwt-user-5", + user_email="admin-edited@example.com", + user_role="internal_user", + ) + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=persisted_row + ) + + result = await get_user_object( + user_id="jwt-user-5", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + user_id_upsert=False, + proxy_logging_obj=None, + user_email="jwt-user-5@example.com", + ) + + assert result is not None + assert result.user_email == "admin-edited@example.com" + + refreshed = await cache.async_get_cache( + key="jwt-user-5", model_type=LiteLLM_UserTable + ) + assert refreshed is not None + assert refreshed.user_email == "admin-edited@example.com" + + @pytest.mark.asyncio async def test_get_user_object_upsert_routes_default_team_to_membership(monkeypatch): """Regression for LIT-4324: a configured default team (list of NewUserRequestTeam diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index affaaa3fbf4..3177fc5ba44 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1989,6 +1989,231 @@ class TestJWTOAuth2Coexistence: assert result.org_id == "validated-org" assert result.user_email == "validated@example.com" + @pytest.mark.asyncio + async def test_mapped_virtual_key_backfills_and_sets_user_email(self): + """ + Regression (LIT-4710): when a JWT resolves straight to an existing + virtual-key mapping (skipping auth_builder), the token's user_email must + still backfill the resolved user and be set on the returned + UserAPIKeyAuth. Before the fix the mapped path never passed the email + through, so user_api_key_user_email stayed null on every request. + """ + jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + general_settings = {"enable_jwt_auth": True} + user_api_key_cache = DualCache() + prisma_client = MagicMock() + jwt_handler = MagicMock() + jwt_handler.is_jwt.return_value = True + jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "mapped-user"}) + jwt_handler.get_user_email = MagicMock(return_value="mapped@example.com") + jwt_handler.get_user_id = MagicMock(return_value="mapped-user") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + user_email_jwt_field="sub", + virtual_key_mapping_cache_ttl=300, + ) + + mapped_key = UserAPIKeyAuth( + token="hashed-mapped-key", + api_key="hashed-mapped-key", + user_id="mapped-user", + user_email=None, + ) + backfilled_user = LiteLLM_UserTable( + user_id="mapped-user", + user_email="mapped@example.com", + user_role="internal_user", + ) + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), + patch( + "litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key", + new_callable=AsyncMock, + return_value=mapped_key, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + new_callable=AsyncMock, + return_value=backfilled_user, + ) as mock_get_user_object, + ): + result = await _user_api_key_auth_builder( + request=mock_request, + api_key=jwt_token, + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "gpt-4o-mini"}, + ) + + assert result.user_id == "mapped-user" + assert result.user_email == "mapped@example.com" + assert ( + mock_get_user_object.call_args_list[0].kwargs["user_email"] + == "mapped@example.com" + ) + + @pytest.mark.asyncio + async def test_mapped_virtual_key_does_not_backfill_mismatched_owner(self): + """ + LIT-4710 security guard: when an admin-created mapping points a JWT at a + virtual key owned by a different user, the JWT principal's email must not + be written onto the mapped key owner's record. Backfill only runs when the + mapped key owner is the JWT principal. + """ + jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + general_settings = {"enable_jwt_auth": True} + user_api_key_cache = DualCache() + prisma_client = MagicMock() + jwt_handler = MagicMock() + jwt_handler.is_jwt.return_value = True + jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "jwt-principal"}) + jwt_handler.get_user_email = MagicMock(return_value="principal@example.com") + jwt_handler.get_user_id = MagicMock(return_value="jwt-principal") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + user_email_jwt_field="sub", + virtual_key_mapping_cache_ttl=300, + ) + + mapped_key = UserAPIKeyAuth( + token="hashed-mapped-key", + api_key="hashed-mapped-key", + user_id="other-owner", + user_email=None, + ) + other_owner = LiteLLM_UserTable( + user_id="other-owner", + user_email=None, + user_role="internal_user", + ) + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), + patch( + "litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key", + new_callable=AsyncMock, + return_value=mapped_key, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + new_callable=AsyncMock, + return_value=other_owner, + ) as mock_get_user_object, + ): + result = await _user_api_key_auth_builder( + request=mock_request, + api_key=jwt_token, + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "gpt-4o-mini"}, + ) + + assert result.user_id == "other-owner" + assert result.user_email is None + assert all( + call.kwargs.get("user_email") != "principal@example.com" + for call in mock_get_user_object.call_args_list + ) + + @pytest.mark.asyncio + async def test_mapped_virtual_key_backfill_failure_does_not_break_auth(self): + """ + LIT-4710 resilience: a mapped-key request served from a valid cached key + must still authenticate when the best-effort email backfill cannot reach + the database, retaining null email rather than failing the request. + """ + jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + general_settings = {"enable_jwt_auth": True} + user_api_key_cache = DualCache() + prisma_client = MagicMock() + jwt_handler = MagicMock() + jwt_handler.is_jwt.return_value = True + jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "mapped-user"}) + jwt_handler.get_user_email = MagicMock(return_value="mapped@example.com") + jwt_handler.get_user_id = MagicMock(return_value="mapped-user") + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + virtual_key_claim_field="sub", + user_email_jwt_field="sub", + virtual_key_mapping_cache_ttl=300, + ) + + mapped_key = UserAPIKeyAuth( + token="hashed-mapped-key", + api_key="hashed-mapped-key", + user_id="mapped-user", + user_email=None, + ) + + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + with ( + patch("litellm.proxy.proxy_server.general_settings", general_settings), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), + patch( + "litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key", + new_callable=AsyncMock, + return_value=mapped_key, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_user_object", + new_callable=AsyncMock, + side_effect=Exception("can't reach database server"), + ), + ): + result = await _user_api_key_auth_builder( + request=mock_request, + api_key=jwt_token, + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "gpt-4o-mini"}, + ) + + assert result.user_id == "mapped-user" + assert result.user_email is None + @pytest.mark.asyncio async def test_routing_override_routes_matching_jwt_to_oauth2(self): """