mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 2ea8126377 into b781d157d7
This commit is contained in:
commit
3516bd6343
2 changed files with 168 additions and 13 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue