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.
This commit is contained in:
Sergey Stepanenko 2026-09-29 11:08:33 +02:00
parent 85dc7cb62e
commit 2ea8126377
2 changed files with 168 additions and 13 deletions

View file

@ -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(

View file

@ -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()