From 2ea8126377e999aff2e4798940f60cf52c5ecd14 Mon Sep 17 00:00:00 2001 From: Sergey Stepanenko Date: Tue, 29 Sep 2026 11:08:33 +0200 Subject: [PATCH] fix(auth): recover concurrent JWT user creation conflicts Two parallel requests carrying the same JWT for a user the proxy has not seen yet both miss the cache and the DB read inside get_user_object with user_id_upsert on, then both call create. The loser fails on the user_id unique constraint, the outer handler turns that into "User doesn't exist in db" and the request is rejected with 401, so some of the very first requests of a new user can fail depending on how its client parallelises. On a unique-constraint conflict the row the winner inserted is now read back from the writer, so read-replica lag cannot hide it, and auth carries on with that row. Default-team enrolment stays on the creation path, so only the request that inserted the row runs it. Any other create error still fails as before. --- litellm/proxy/auth/auth_checks.py | 42 +++++-- tests/unit/proxy/auth/test_auth_checks.py | 139 ++++++++++++++++++++++ 2 files changed, 168 insertions(+), 13 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 51cc70c010b..a71268f047a 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -15,7 +15,7 @@ import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias +from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias, cast from fastapi import HTTPException, Request, status from pydantic import BaseModel, TypeAdapter, ValidationError @@ -299,6 +299,13 @@ def _user_table(repo: _PrismaTableHolder[_PrismaUserRow]) -> _PrismaAuthTable[_P return _DeadlineBoundedTable(repo.table, "user") +def _user_writer_table(prisma_client: PrismaClient) -> _PrismaAuthTable[_PrismaUserRow]: + """The user table on the writer, so a row another request just inserted is not hidden by replica lag.""" + writer: Final = prisma_client.writer_db + table: Final = cast("_PrismaAuthTable[_PrismaUserRow]", writer.litellm_usertable) # cast-ok: untyped wrapper + return _DeadlineBoundedTable(table, "user") + + class _VectorStorePermissionsRow(Protocol): @property def vector_stores(self) -> Sequence[str] | None: ... @@ -2680,20 +2687,29 @@ async def get_user_object( budget_duration=new_user_params["budget_duration"] ) - response = await _user_table(UserRepository(prisma_client)).create( - data=new_user_params, - include={"organization_memberships": True}, - ) + from prisma.errors import UniqueViolationError - default_teams: Final = check_if_default_team_set() - if default_teams: - await add_new_user_to_default_team( - user_id=user_id, - user_email=user_email, - user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), - teams=default_teams, - prisma_client=prisma_client, + try: + response = await _user_table(UserRepository(prisma_client)).create( + data=new_user_params, + include={"organization_memberships": True}, ) + except UniqueViolationError: + response = await _user_writer_table(prisma_client).find_unique( + where={"user_id": user_id}, include={"organization_memberships": True} + ) + if response is None: + raise + else: + default_teams: Final = check_if_default_team_set() + if default_teams: + await add_new_user_to_default_team( + user_id=user_id, + user_email=user_email, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + teams=default_teams, + prisma_client=prisma_client, + ) else: if should_check_db: _update_last_db_access_time( diff --git a/tests/unit/proxy/auth/test_auth_checks.py b/tests/unit/proxy/auth/test_auth_checks.py index 2538556d3b5..cb85e2c14ca 100644 --- a/tests/unit/proxy/auth/test_auth_checks.py +++ b/tests/unit/proxy/auth/test_auth_checks.py @@ -3,12 +3,15 @@ import sys, os, asyncio, time, random, uuid import traceback +from typing import Final +from unittest.mock import AsyncMock, MagicMock from dotenv import load_dotenv load_dotenv() import pytest, litellm import httpx +from prisma.errors import UniqueViolationError from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import get_end_user_object from litellm.caching.caching import DualCache @@ -25,6 +28,7 @@ from litellm.proxy.auth.auth_checks import ( can_team_access_model, _virtual_key_soft_budget_check, _team_soft_budget_check, + get_user_object, ) from litellm.proxy.utils import ProxyLogging from litellm.proxy.utils import CallInfo @@ -1491,3 +1495,138 @@ async def test_key_access_group_grants_model_when_get_access_object_raises(): finally: for p in patches: p.stop() + + +@pytest.fixture +def user_upsert_db(monkeypatch: pytest.MonkeyPatch) -> MagicMock: + """A database with no row for the user on the routed side, so an upsert reaches create.""" + monkeypatch.setattr(litellm, "default_internal_user_params", None) + db: Final = MagicMock() + db.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + db.db.litellm_usertable.find_first = AsyncMock(return_value=None) + return db + + +def _user_id_conflict() -> UniqueViolationError: + return UniqueViolationError({"message": "Unique constraint failed on the fields: (`user_id`)"}) + + +async def _user_upsert(user_id: str, db: MagicMock, cache: UserApiKeyCache) -> LiteLLM_UserTable | None: + return await get_user_object( + user_id=user_id, + prisma_client=db, + user_api_key_cache=cache, + user_id_upsert=True, + proxy_logging_obj=None, + user_email=f"{user_id}@example.com", + ) + + +@pytest.mark.asyncio +async def test_get_user_object_upsert_adopts_the_row_a_concurrent_request_created(user_upsert_db: MagicMock) -> None: + """Two requests for a new user both miss the cache and the read, then both call create; + the loser must adopt the winner's row instead of failing auth.""" + created: Final = LiteLLM_UserTable(user_id="jwt-new-user", user_email="jwt-new-user@example.com") + first_at_create: Final = asyncio.Event() + both_at_create: Final = asyncio.Event() + winner_inserted: Final = asyncio.Event() + + async def create(**_kwargs: object) -> LiteLLM_UserTable: + if first_at_create.is_set(): + both_at_create.set() + first_at_create.set() + await asyncio.wait_for(both_at_create.wait(), timeout=5) + if winner_inserted.is_set(): + raise _user_id_conflict() + winner_inserted.set() + return created + + user_upsert_db.db.litellm_usertable.create = AsyncMock(side_effect=create) + user_upsert_db.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=created) + + results: Final = await asyncio.wait_for( + asyncio.gather(*(_user_upsert("jwt-new-user", user_upsert_db, UserApiKeyCache()) for _ in range(2))), timeout=10 + ) + + assert [user.user_id for user in results if user is not None] == ["jwt-new-user", "jwt-new-user"] + assert user_upsert_db.db.litellm_usertable.create.await_count == 2, "both requests reached create" + user_upsert_db.writer_db.litellm_usertable.find_unique.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_get_user_object_upsert_loser_reads_the_winner_row_from_the_writer_and_caches_it( + user_upsert_db: MagicMock, +) -> None: + winner_row: Final = LiteLLM_UserTable( + user_id="jwt-loser", user_email="jwt-loser@example.com", user_role="internal_user", max_budget=12.5, models=["m1"] + ) + user_upsert_db.db.litellm_usertable.create = AsyncMock(side_effect=_user_id_conflict()) + user_upsert_db.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=winner_row) + cache: Final = UserApiKeyCache() + + user: Final = await _user_upsert("jwt-loser", user_upsert_db, cache) + + assert user == winner_row, "the loser must carry the winner's limits, role and models" + user_upsert_db.writer_db.litellm_usertable.find_unique.assert_awaited_once_with( + where={"user_id": "jwt-loser"}, include={"organization_memberships": True} + ) + assert await cache.async_get_cache(key="jwt-loser", model_type=LiteLLM_UserTable) == winner_row + + +@pytest.mark.asyncio +async def test_get_user_object_upsert_conflict_without_a_readable_row_still_fails(user_upsert_db: MagicMock) -> None: + conflict: Final = _user_id_conflict() + user_upsert_db.db.litellm_usertable.create = AsyncMock(side_effect=conflict) + user_upsert_db.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=None) + + with pytest.raises(ValueError, match="jwt-ghost") as excinfo: + await _user_upsert("jwt-ghost", user_upsert_db, UserApiKeyCache()) + + user_upsert_db.writer_db.litellm_usertable.find_unique.assert_awaited_once() + assert excinfo.value.__context__ is conflict, "the original conflict surfaces, not a silent None" + + +@pytest.mark.asyncio +async def test_get_user_object_upsert_does_not_recover_from_an_unrelated_create_error( + user_upsert_db: MagicMock, +) -> None: + user_upsert_db.db.litellm_usertable.create = AsyncMock(side_effect=RuntimeError("column does not exist")) + user_upsert_db.writer_db.litellm_usertable.find_unique = AsyncMock() + + with pytest.raises(ValueError, match="column does not exist"): + await _user_upsert("jwt-unlucky", user_upsert_db, UserApiKeyCache()) + + user_upsert_db.writer_db.litellm_usertable.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_get_user_object_upsert_enrolls_the_user_it_created_in_the_default_teams( + monkeypatch: pytest.MonkeyPatch, user_upsert_db: MagicMock +) -> None: + monkeypatch.setattr(litellm, "default_internal_user_params", {"teams": ["team-default"]}) + user_upsert_db.db.litellm_usertable.create = AsyncMock(return_value=LiteLLM_UserTable(user_id="jwt-first", user_email="jwt-first@example.com")) + enroll: Final = AsyncMock() + monkeypatch.setattr("litellm.proxy.management_endpoints.internal_user_endpoints.add_new_user_to_default_team", enroll) + + await _user_upsert("jwt-first", user_upsert_db, UserApiKeyCache()) + + enroll.assert_awaited_once() + assert enroll.await_args is not None + assert enroll.await_args.kwargs["user_id"] == "jwt-first" + assert enroll.await_args.kwargs["teams"] == ["team-default"] + + +@pytest.mark.asyncio +async def test_get_user_object_upsert_recovery_does_not_enroll_the_user_again( + monkeypatch: pytest.MonkeyPatch, user_upsert_db: MagicMock +) -> None: + monkeypatch.setattr(litellm, "default_internal_user_params", {"teams": ["team-default"]}) + user_upsert_db.db.litellm_usertable.create = AsyncMock(side_effect=_user_id_conflict()) + user_upsert_db.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=LiteLLM_UserTable(user_id="jwt-second", user_email="jwt-second@example.com")) + enroll: Final = AsyncMock() + monkeypatch.setattr("litellm.proxy.management_endpoints.internal_user_endpoints.add_new_user_to_default_team", enroll) + + user: Final = await _user_upsert("jwt-second", user_upsert_db, UserApiKeyCache()) + + assert user is not None and user.user_id == "jwt-second" + enroll.assert_not_awaited()