mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
test: expose replica_db on remaining prisma test doubles and unshadow body helper in mcp env var test
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e1f9382378
commit
ebf89e4ae6
13 changed files with 146 additions and 11 deletions
|
|
@ -195,7 +195,7 @@ def test_malformed_bodies_missing_users_and_foreign_servers_are_rejected(gateway
|
|||
path: Final = f"/v1/mcp/server/{identity}/user-env-vars"
|
||||
payload: Final[dict[str, JsonValue]] = {"values": {TOKEN: "x"}}
|
||||
malformed: Final[tuple[dict[str, JsonValue], ...]] = ({"values": {TOKEN: 7}}, {"values": ["a"]}, {})
|
||||
assert [gateway.request("POST", path, body, key=key).status_code for body in malformed] == [422, 422, 422]
|
||||
assert [gateway.request("POST", path, bad, key=key).status_code for bad in malformed] == [422, 422, 422]
|
||||
assert set_names(env_status(gateway, key, identity)) == {TOKEN: False}
|
||||
assert [gateway.client.request(method, path, json=payload).status_code for method in METHODS] == [401, 401, 401]
|
||||
no_user: Final = tuple(gateway.request(method, path, payload, key=userless) for method in METHODS)
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ def _build_prisma_mock(db_batch_object=None):
|
|||
mock.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
|
||||
mock.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=db_batch_object)
|
||||
mock.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
mock.replica_db = mock.db
|
||||
return mock
|
||||
|
||||
|
||||
|
|
@ -376,6 +377,7 @@ def _in_memory_managed_files():
|
|||
prisma = MagicMock()
|
||||
prisma.db.litellm_managedobjecttable = table
|
||||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
||||
prisma.replica_db = prisma.db
|
||||
|
||||
cache = MagicMock()
|
||||
cache.async_set_cache = AsyncMock()
|
||||
|
|
|
|||
|
|
@ -153,6 +153,7 @@ class TestCloudZeroHourlyExport:
|
|||
fake_db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
|
||||
fake_db.litellm_usertable.find_many = AsyncMock(return_value=[])
|
||||
fake_client.db = fake_db
|
||||
fake_client.replica_db = fake_db
|
||||
mock_prisma_client_getter.return_value = fake_client
|
||||
|
||||
mock_datetime.now.return_value = datetime(2025, 11, 1, 12, 0, 1)
|
||||
|
|
@ -182,6 +183,7 @@ class TestLiteLLMDatabaseUsageData:
|
|||
|
||||
fake_client = MagicMock()
|
||||
fake_client.db.query_raw = AsyncMock(side_effect=query_raw)
|
||||
fake_client.replica_db = fake_client.db
|
||||
db = LiteLLMDatabase()
|
||||
monkeypatch.setattr(db, "_ensure_prisma_client", lambda: fake_client)
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,8 @@ from litellm.integrations.cloudzero.database import LiteLLMDatabase
|
|||
def _setup_db(monkeypatch: pytest.MonkeyPatch, query_return):
|
||||
"""Return a database instance with prisma client mocked out."""
|
||||
query_mock = AsyncMock(return_value=query_return)
|
||||
mock_client = SimpleNamespace(db=SimpleNamespace(query_raw=query_mock))
|
||||
mock_db = SimpleNamespace(query_raw=query_mock)
|
||||
mock_client = SimpleNamespace(db=mock_db, replica_db=mock_db)
|
||||
db = LiteLLMDatabase()
|
||||
monkeypatch.setattr(db, "_ensure_prisma_client", lambda: mock_client)
|
||||
return db, query_mock
|
||||
|
|
|
|||
|
|
@ -13,7 +13,8 @@ from litellm.integrations.focus.database import FocusLiteLLMDatabase
|
|||
def _setup_db(monkeypatch: pytest.MonkeyPatch, query_return):
|
||||
"""Create a database instance with a stubbed prisma client."""
|
||||
query_mock = AsyncMock(return_value=query_return)
|
||||
mock_client = SimpleNamespace(db=SimpleNamespace(query_raw=query_mock))
|
||||
mock_db = SimpleNamespace(query_raw=query_mock)
|
||||
mock_client = SimpleNamespace(db=mock_db, replica_db=mock_db)
|
||||
db = FocusLiteLLMDatabase()
|
||||
monkeypatch.setattr(db, "_ensure_prisma_client", lambda: mock_client)
|
||||
return db, query_mock
|
||||
|
|
@ -101,7 +102,8 @@ async def test_should_build_frame_from_rows_recovered_for_double_hashed_keys(mon
|
|||
return [{"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": None}]
|
||||
return [joined_row, dirty_row]
|
||||
|
||||
mock_client = SimpleNamespace(db=SimpleNamespace(query_raw=AsyncMock(side_effect=query_raw)))
|
||||
mock_db = SimpleNamespace(query_raw=AsyncMock(side_effect=query_raw))
|
||||
mock_client = SimpleNamespace(db=mock_db, replica_db=mock_db)
|
||||
db = FocusLiteLLMDatabase()
|
||||
monkeypatch.setattr(db, "_ensure_prisma_client", lambda: mock_client)
|
||||
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ def _make_prisma_with_existing(row):
|
|||
"""Build a MagicMock prisma_client whose user-credentials table returns ``row``
|
||||
for find_unique and behaves async-correctly for upsert/find_many."""
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=row)
|
||||
prisma.db.litellm_mcpusercredentials.upsert = AsyncMock()
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[])
|
||||
|
|
@ -195,6 +196,7 @@ async def test_purge_user_oauth_credentials_for_server_invalidates_each_user():
|
|||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice"), _oauth_row("bob")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
|
||||
|
|
@ -233,6 +235,7 @@ async def test_list_server_user_credentials_types_each_row_without_leaking_the_s
|
|||
byok_row = _byok_row("carol")
|
||||
byok_row.updated_at = datetime(2026, 2, 1, tzinfo=timezone.utc)
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[oauth_row, byok_row])
|
||||
|
||||
items = await list_server_user_credentials(prisma, "srv-1")
|
||||
|
|
@ -267,6 +270,7 @@ async def test_purge_user_oauth_credentials_for_server_spares_byok_rows():
|
|||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_byok_row("carol"), _oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=1)
|
||||
|
||||
|
|
@ -290,6 +294,7 @@ async def test_purge_user_oauth_credentials_for_server_all_byok_is_noop():
|
|||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_byok_row("carol"), _byok_row("dave")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
|
||||
|
||||
|
|
@ -309,6 +314,7 @@ async def test_purge_user_oauth_credentials_for_server_defaults_to_manager_inval
|
|||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=1)
|
||||
|
||||
|
|
@ -331,6 +337,7 @@ async def test_purge_user_oauth_credentials_for_server_logs_raced_rows(monkeypat
|
|||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=0)
|
||||
warning = MagicMock()
|
||||
|
|
@ -350,6 +357,7 @@ async def test_delete_mcp_server_invalidates_cached_tokens_for_enumerated_users(
|
|||
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=MagicMock(server_id="srv-1"))
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[_oauth_row("alice"), _byok_row("bob")])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock(return_value=2)
|
||||
|
|
@ -371,6 +379,7 @@ async def test_delete_mcp_server_returns_none_without_cleanup_when_server_missin
|
|||
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=None)
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock()
|
||||
|
||||
|
|
@ -385,6 +394,7 @@ async def test_purge_user_oauth_credentials_for_server_noop_when_empty():
|
|||
from litellm.proxy._experimental.mcp_server.db import purge_user_oauth_credentials_for_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
|
||||
|
||||
|
|
@ -464,7 +474,8 @@ class _MapTable:
|
|||
@pytest.mark.parametrize("quoted", [False, True])
|
||||
async def test_secret_maps_create_update_round_trip(map_algorithm: str, field: str, quoted: bool) -> None:
|
||||
table: Final = _MapTable(quoted=quoted)
|
||||
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
|
||||
tables: Final = SimpleNamespace(litellm_mcpservertable=table)
|
||||
prisma: Final = SimpleNamespace(db=tables, replica_db=tables)
|
||||
original: Final = {"TOKEN": " sensitive-secret\n", "PREFIX": "v2:gcm:literal", "TEMPLATE": "Bearer ${TOKEN}"}
|
||||
create: Final = NewMCPServerRequest.model_validate({
|
||||
"server_id": "srv-map", "transport": "http", "url": "https://up.example.com/mcp", field: original,
|
||||
|
|
@ -536,9 +547,10 @@ async def test_secret_map_rotation_migrates_rekeys_and_preserves_corrupt(
|
|||
{"server_id": "legacy", field: json.dumps(values), other: "{}"},
|
||||
{"server_id": "encrypted", field: old, other: None},
|
||||
)
|
||||
prisma: Final = SimpleNamespace(db=SimpleNamespace(
|
||||
tables: Final = SimpleNamespace(
|
||||
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[]))
|
||||
))
|
||||
)
|
||||
prisma: Final = SimpleNamespace(db=tables, replica_db=tables)
|
||||
await rotate_mcp_server_credentials_master_key(prisma, touched_by="test", new_master_key="rotated-map-key")
|
||||
assert table.rows["broken"][field] == corrupt
|
||||
assert table.rows["legacy"][other] == "{}" and table.rows["encrypted"][other] is None
|
||||
|
|
@ -565,7 +577,8 @@ async def test_bulk_reads_isolate_corrupt_secret_maps(reader, field, map_algorit
|
|||
]
|
||||
snapshot = [row.model_dump() for row in rows]
|
||||
table = SimpleNamespace(find_many=AsyncMock(return_value=rows))
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
|
||||
tables = SimpleNamespace(litellm_mcpservertable=table)
|
||||
prisma = SimpleNamespace(db=tables, replica_db=tables)
|
||||
result = await reader(prisma, ["broken", "healthy"]) if reader is get_mcp_servers else await reader(prisma)
|
||||
items = result.items if reader is get_mcp_submissions else result
|
||||
assert [row.server_id for row in items] == ["healthy"]
|
||||
|
|
@ -584,7 +597,8 @@ async def test_bulk_reads_do_not_swallow_unrelated_validation_errors(reader):
|
|||
|
||||
row = _prisma_map_row({"server_id": "invalid", "transport": "unsupported"})
|
||||
table = SimpleNamespace(find_many=AsyncMock(return_value=[row]))
|
||||
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
|
||||
tables = SimpleNamespace(litellm_mcpservertable=table)
|
||||
prisma = SimpleNamespace(db=tables, replica_db=tables)
|
||||
request = reader(prisma, ["invalid"]) if reader is get_mcp_servers else reader(prisma)
|
||||
with pytest.raises(ValidationError, match="transport"):
|
||||
await request
|
||||
|
|
@ -1316,6 +1330,7 @@ async def test_rotate_user_env_vars_re_encrypts_with_new_key(monkeypatch):
|
|||
encrypted_old = encrypt_value_helper(json.dumps(values))
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[_env_var_row(encrypted_old)])
|
||||
prisma.db.litellm_mcpuserenvvars.update = AsyncMock()
|
||||
|
||||
|
|
@ -1343,6 +1358,7 @@ async def test_rotate_user_env_vars_skips_undecryptable_rows():
|
|||
bad = _env_var_row("!!! not encrypted !!!", server_id="srv-corrupt")
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[bad, good])
|
||||
prisma.db.litellm_mcpuserenvvars.update = AsyncMock()
|
||||
|
||||
|
|
@ -1501,6 +1517,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch):
|
|||
monkeypatch.setattr(enc, "_get_salt_key", lambda: key_old)
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(server_id="config_faros", credentials=blob_old)]
|
||||
|
|
@ -1528,6 +1545,7 @@ async def test_delete_mcp_server_cleans_oauth_client_store():
|
|||
from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.replica_db = prisma.db
|
||||
prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=SimpleNamespace(server_id="s1"))
|
||||
prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[])
|
||||
prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
|
||||
|
|
|
|||
|
|
@ -798,6 +798,7 @@ async def test_default_internal_user_params_with_get_user_object(monkeypatch):
|
|||
mock_prisma_client = MagicMock()
|
||||
mock_db = AsyncMock()
|
||||
mock_prisma_client.db = mock_db
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
# Set up the user creation mock - create a complete user model that can be converted to a dict
|
||||
mock_user = MagicMock()
|
||||
|
|
@ -863,6 +864,7 @@ async def test_get_user_object_upsert_sets_budget_reset_at(monkeypatch, has_budg
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db = AsyncMock()
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.create = AsyncMock(return_value=MagicMock(organization_memberships=[]))
|
||||
|
|
@ -897,6 +899,7 @@ async def test_get_user_object_upsert_sets_budget_reset_at(monkeypatch, has_budg
|
|||
def _user_read_raising(error: Exception) -> tuple[MagicMock, MagicMock]:
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=error)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
cache = MagicMock()
|
||||
cache.async_get_cache = AsyncMock(return_value=None)
|
||||
cache.async_set_cache = AsyncMock()
|
||||
|
|
@ -982,6 +985,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
|
|||
db_row = LiteLLM_UserTable(user_id=user_id, user_email=None, user_role="internal_user")
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
result = await get_user_object(
|
||||
user_id=user_id,
|
||||
|
|
@ -1003,6 +1007,7 @@ async def test_get_user_object_upsert_includes_user_email():
|
|||
mock_prisma_client = MagicMock()
|
||||
mock_db = AsyncMock()
|
||||
mock_prisma_client.db = mock_db
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
# Set up the user creation mock
|
||||
mock_user = MagicMock()
|
||||
|
|
@ -1070,6 +1075,7 @@ async def test_get_user_object_backfills_null_email_from_cache_hit():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_UserTable(
|
||||
user_id="jwt-user-1",
|
||||
|
|
@ -1116,6 +1122,7 @@ async def test_get_user_object_backfills_null_email_from_db_read():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=[db_row, backfilled_row])
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1)
|
||||
|
||||
|
|
@ -1155,6 +1162,7 @@ async def test_get_user_object_does_not_overwrite_existing_email():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
result = await get_user_object(
|
||||
user_id="jwt-user-2",
|
||||
|
|
@ -1188,6 +1196,7 @@ async def test_get_user_object_backfill_race_prefers_db_email():
|
|||
)
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=winner_row)
|
||||
|
||||
result = await get_user_object(
|
||||
|
|
@ -1227,6 +1236,7 @@ async def test_get_user_object_backfill_caches_persisted_email_not_proposed():
|
|||
)
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=persisted_row)
|
||||
|
||||
result = await get_user_object(
|
||||
|
|
@ -1260,6 +1270,7 @@ async def test_get_user_object_upsert_routes_default_team_to_membership(monkeypa
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db = AsyncMock()
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
||||
|
||||
|
|
@ -1329,6 +1340,7 @@ async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeyp
|
|||
mock_prisma_client = MagicMock()
|
||||
mock_db = AsyncMock()
|
||||
mock_prisma_client.db = mock_db
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique.return_value = None
|
||||
|
||||
# Define what our mocked `new_team` function should return
|
||||
|
|
@ -1363,6 +1375,7 @@ async def test_get_team_db_check_does_not_call_new_team_if_exists(mock_new_team,
|
|||
mock_prisma_client = MagicMock()
|
||||
mock_db = AsyncMock()
|
||||
mock_prisma_client.db = mock_db
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique.return_value = MagicMock()
|
||||
|
||||
team_id_to_find = "existing-jwt-team"
|
||||
|
|
@ -1419,6 +1432,7 @@ async def test_vector_store_access_check_skips_db_lookup_when_no_vector_stores_r
|
|||
mock_prisma_client = MagicMock()
|
||||
find_unique = AsyncMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = find_unique
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_vector_store_registry = MagicMock()
|
||||
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = []
|
||||
|
|
@ -1515,6 +1529,7 @@ async def test_vector_store_access_check_with_permissions():
|
|||
mock_permissions = MagicMock()
|
||||
mock_permissions.vector_stores = ["store-1", "store-2"]
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=mock_permissions)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_vector_store_registry = MagicMock()
|
||||
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-1"]
|
||||
|
|
@ -1561,6 +1576,7 @@ async def test_vector_store_access_check_with_team_permissions():
|
|||
team_permissions = MagicMock()
|
||||
team_permissions.vector_stores = ["team-store-allowed"]
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_vector_store_registry = MagicMock()
|
||||
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["team-store-allowed"]
|
||||
|
|
@ -2257,6 +2273,7 @@ async def test_get_tag_objects_batch():
|
|||
|
||||
# Mock DB to return all uncached tags in ONE query
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(return_value=[uncached_tag_1, uncached_tag_2, uncached_tag_3])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
||||
# Call batch fetch
|
||||
tag_objects = await get_tag_objects_batch(
|
||||
|
|
@ -2359,6 +2376,7 @@ async def test_get_tag_objects_batch_never_queries_db_for_unregistered_tags():
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(return_value=[_tag_registry_row("some-other-tag")])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
first = await get_tag_objects_batch(
|
||||
|
|
@ -2400,6 +2418,7 @@ async def test_get_tag_objects_batch_fetches_only_registered_uncached_tags():
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
||||
tag_objects = await get_tag_objects_batch(
|
||||
tag_names=["cached-tag", "registered-tag", "unregistered-tag"],
|
||||
|
|
@ -2422,6 +2441,7 @@ async def test_get_tag_objects_batch_caches_empty_registry():
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
assert (
|
||||
|
|
@ -2468,6 +2488,7 @@ async def test_get_tag_objects_batch_registry_db_error_negative_caches_and_keeps
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
cache = _TtlRecordingCache()
|
||||
|
||||
first = await get_tag_objects_batch(
|
||||
|
|
@ -2516,6 +2537,7 @@ async def test_tag_registry_load_is_single_flighted_across_concurrent_requests()
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
results = await asyncio.gather(
|
||||
|
|
@ -2547,6 +2569,7 @@ async def test_get_tag_objects_batch_oversized_registry_falls_back_and_stops_ref
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
first = await get_tag_objects_batch(
|
||||
|
|
@ -2584,6 +2607,7 @@ async def test_tag_max_budget_check_still_enforces_registered_tag_over_budget():
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
|
||||
if counter_key == "spend:tag:paid-tag":
|
||||
|
|
@ -2618,6 +2642,7 @@ async def test_get_team_object_raises_404_when_not_found():
|
|||
mock_prisma_client = MagicMock()
|
||||
mock_db = AsyncMock()
|
||||
mock_prisma_client.db = mock_db
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
|
|
@ -2641,6 +2666,7 @@ def _mock_prisma_for_team_lookup(find_unique):
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = find_unique
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
return mock_prisma_client
|
||||
|
||||
|
||||
|
|
@ -3583,6 +3609,7 @@ async def test_get_fuzzy_user_object_case_insensitive_email():
|
|||
# Setup mock Prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_usertable = MagicMock()
|
||||
|
||||
# Mock user data with mixed case email
|
||||
|
|
@ -4263,6 +4290,7 @@ async def test_team_member_budget_check_falls_back_to_team_default_budget_id():
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_budget_row)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
|
||||
if counter_key == "spend:team_member:test-user:test-team":
|
||||
|
|
@ -4293,6 +4321,7 @@ async def test_team_member_budget_check_falls_back_to_team_default_budget_id():
|
|||
|
||||
# First call did perform the fallback DB lookup.
|
||||
prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once()
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
# Second call hits the cached budget row, no additional prisma read.
|
||||
prisma_client.db.litellm_budgettable.find_unique.reset_mock()
|
||||
|
|
@ -4356,6 +4385,7 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_budget_row)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
mocked_spend = 70.0
|
||||
|
||||
|
|
@ -4383,6 +4413,7 @@ async def test_team_member_budget_check_per_member_override_wins_over_team_defau
|
|||
)
|
||||
|
||||
prisma_client.db.litellm_budgettable.find_unique.assert_not_awaited()
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
# 2. Now push spend above the per-member cap ($200). Must raise with
|
||||
# max_budget=200 to prove the per-member cap is the value being
|
||||
|
|
@ -4444,6 +4475,7 @@ async def test_team_member_budget_check_null_clone_falls_back_to_team_default():
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_default_row)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
|
||||
if counter_key == "spend:team_member:test-user:test-team":
|
||||
|
|
@ -4471,6 +4503,7 @@ async def test_team_member_budget_check_null_clone_falls_back_to_team_default():
|
|||
assert exc_info.value.current_cost == 500.0
|
||||
assert exc_info.value.max_budget == 65.0
|
||||
prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once()
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -4507,6 +4540,7 @@ async def test_team_member_budget_check_null_clone_with_null_default_skips_enfor
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_default_row)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
|
||||
if counter_key == "spend:team_member:test-user:test-team":
|
||||
|
|
@ -4570,6 +4604,7 @@ async def test_team_member_budget_check_zero_team_default_treated_as_no_cap():
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=fake_default_row)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
|
||||
if counter_key == "spend:team_member:test-user:test-team":
|
||||
|
|
@ -4628,6 +4663,7 @@ async def test_team_member_budget_check_zero_per_member_row_still_blocks():
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
|
||||
if counter_key == "spend:team_member:test-user:test-team":
|
||||
|
|
@ -6167,6 +6203,7 @@ async def test_get_org_object_for_request_serves_last_known_org_through_db_outag
|
|||
prisma_client.db.litellm_organizationtable.find_unique = AsyncMock(
|
||||
side_effect=[db_outage] if warmed_by_auth_prefetch else [org_row, db_outage]
|
||||
)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
user_api_key_cache = UserApiKeyCache()
|
||||
if warmed_by_auth_prefetch:
|
||||
await user_api_key_cache.async_set_cache(
|
||||
|
|
@ -6348,6 +6385,7 @@ async def test_get_default_end_user_budget_db_fetch_returns_validated_budget(mon
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -6383,6 +6421,7 @@ async def test_get_team_member_default_budget_caches_json_safe_payload():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
class _JsonOnlyRedis:
|
||||
"""Stands in for RedisCache, which serializes with a bare json.dumps()."""
|
||||
|
|
@ -6421,6 +6460,7 @@ async def test_get_team_member_default_budget_caches_json_safe_payload():
|
|||
assert isinstance(cached, LiteLLM_BudgetTable)
|
||||
assert cached.max_budget == 25.0
|
||||
mock_prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once()
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -6432,6 +6472,7 @@ async def test_get_end_user_object_db_fetch_returns_validated_end_user():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -6495,6 +6536,7 @@ async def test_get_end_user_object_never_queries_db_for_unrestricted_end_users(
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[_end_user_registry_row("eu-blocked")])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
|
|
@ -6536,6 +6578,7 @@ async def test_get_end_user_object_still_fetches_restricted_end_user(end_user_re
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[_end_user_registry_row("eu-blocked")])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
||||
return_value=_end_user_db_row("eu-blocked", blocked=True)
|
||||
)
|
||||
|
|
@ -6569,6 +6612,7 @@ async def test_get_end_user_object_caches_empty_restricted_registry(end_user_reg
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
|
|
@ -6612,6 +6656,7 @@ async def test_get_end_user_object_registry_db_error_negative_caches_and_keeps_p
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(side_effect=Exception("registry query failed"))
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
||||
side_effect=lambda **kwargs: _end_user_db_row(kwargs["where"]["user_id"], blocked=True)
|
||||
)
|
||||
|
|
@ -6659,6 +6704,7 @@ async def test_registry_db_error_is_logged_at_warning(end_user_registry_skip_ena
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(side_effect=Exception("registry query failed"))
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-1", blocked=True))
|
||||
|
||||
with patch("litellm.proxy.auth.auth_checks.verbose_proxy_logger") as mock_logger:
|
||||
|
|
@ -6693,6 +6739,7 @@ async def test_end_user_registry_load_is_single_flighted_across_concurrent_reque
|
|||
return [_end_user_registry_row("eu-blocked")]
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(side_effect=fake_find_many)
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
||||
cache = UserApiKeyCache()
|
||||
|
|
@ -6723,6 +6770,7 @@ async def test_get_end_user_object_oversized_registry_falls_back_and_stops_refet
|
|||
oversized = [_end_user_registry_row(f"eu-{index}") for index in range(END_USER_RESTRICTED_REGISTRY_MAX_SIZE + 1)]
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=oversized)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
||||
side_effect=lambda **kwargs: _end_user_db_row(kwargs["where"]["user_id"], blocked=True)
|
||||
)
|
||||
|
|
@ -6766,6 +6814,7 @@ async def test_get_end_user_object_default_budget_gate_keeps_fetching_unrestrict
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
||||
mock_prisma.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row)
|
||||
|
||||
|
|
@ -6796,6 +6845,7 @@ async def test_get_end_user_object_token_budget_gate_keeps_fetching_unrestricted
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1", spend=100.0))
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
|
|
@ -6841,6 +6891,7 @@ async def test_get_end_user_object_key_default_budget_beats_global_default_witho
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-shared"))
|
||||
mock_prisma.db.litellm_budgettable.find_unique = _budget_lookup_by_id(
|
||||
{"global-eu-budget": 100.0, "svc-a-budget": 0.5, "svc-b-budget": 7.0}
|
||||
|
|
@ -6885,6 +6936,7 @@ async def test_get_end_user_object_cached_row_does_not_carry_another_keys_defaul
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-shared"))
|
||||
mock_prisma.db.litellm_budgettable.find_unique = _budget_lookup_by_id({"svc-a-budget": 0.5})
|
||||
cache = UserApiKeyCache()
|
||||
|
|
@ -6919,6 +6971,7 @@ async def test_get_end_user_object_caches_row_with_global_default_but_never_a_ke
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-cached"))
|
||||
mock_prisma.db.litellm_budgettable.find_unique = _budget_lookup_by_id({"svc-a-budget": 0.5, "global-budget": 7.0})
|
||||
cache = UserApiKeyCache()
|
||||
|
|
@ -6948,6 +7001,7 @@ async def test_get_end_user_object_key_default_budget_loads_unrestricted_row_wit
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1", spend=3.0))
|
||||
mock_prisma.db.litellm_budgettable.find_unique = _budget_lookup_by_id({"svc-a-budget": 2.0})
|
||||
|
||||
|
|
@ -6980,6 +7034,7 @@ async def test_get_end_user_object_explicit_end_user_budget_beats_key_default(mo
|
|||
litellm_budget_table={"budget_id": "vip-budget", "max_budget": 500.0},
|
||||
)
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_budgettable.find_unique = _budget_lookup_by_id({"svc-a-budget": 0.5})
|
||||
|
||||
result = await get_end_user_object(
|
||||
|
|
@ -7002,6 +7057,7 @@ async def test_resolve_default_end_user_budget_falls_back_to_global_when_key_bud
|
|||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_budgettable.find_unique = _budget_lookup_by_id({"global-eu-budget": 100.0})
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
||||
resolved = await resolve_default_end_user_budget(
|
||||
prisma_client=mock_prisma,
|
||||
|
|
@ -7028,6 +7084,7 @@ async def test_end_user_id_validation_gate_still_resolves_unrestricted_end_users
|
|||
monkeypatch.setattr(litellm, "validate_end_user_id_in_db", True)
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-known-1"))
|
||||
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
||||
|
|
@ -7053,6 +7110,7 @@ async def test_get_team_membership_db_fetch_returns_validated_membership():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -7087,6 +7145,7 @@ async def test_get_team_membership_negative_caches_a_missing_row():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
|
|
@ -7124,6 +7183,7 @@ async def test_get_team_membership_reads_sentinel_as_no_membership_not_a_model()
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
result = await get_team_membership(
|
||||
user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache
|
||||
|
|
@ -7149,6 +7209,7 @@ async def test_get_team_membership_coalesces_parallel_db_fetches():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_slow_find_unique)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
async def _load():
|
||||
|
|
@ -7170,6 +7231,7 @@ async def test_get_team_membership_coalesces_parallel_db_fetches():
|
|||
assert results[0].user_id == "u-parallel"
|
||||
assert results[1].user_id == "u-parallel"
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -7194,6 +7256,7 @@ async def test_get_team_membership_invalidation_waits_for_in_flight_load_then_ev
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_find_unique)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
cache = UserApiKeyCache()
|
||||
_key = team_membership_reservation_cache_key(user_id="u-inv", team_id="t-inv")
|
||||
|
||||
|
|
@ -7245,6 +7308,7 @@ async def test_get_team_membership_invalidation_during_cache_write_evicts_stale_
|
|||
row.dict = lambda: {"user_id": "u-w", "team_id": "t-w", "spend": 1.0, "budget_id": "budget-old"}
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
cache = _SlowWriteCache()
|
||||
|
||||
stale = asyncio.create_task(
|
||||
|
|
@ -7362,6 +7426,7 @@ async def test_get_team_membership_db_error_surfaces_and_retries_next_call():
|
|||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(
|
||||
side_effect=[RuntimeError("db down"), membership_row]
|
||||
)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
with pytest.raises(RuntimeError, match="db down"):
|
||||
|
|
@ -7394,6 +7459,8 @@ class _UnreachableMembershipPrisma:
|
|||
async def find_unique(where: dict[str, dict[str, str]], include: dict[str, bool]) -> None:
|
||||
raise httpx.ConnectError("All connection attempts failed")
|
||||
|
||||
replica_db = db
|
||||
|
||||
|
||||
def _restricted_member_check_deps() -> dict[str, object]:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
|
@ -7445,6 +7512,7 @@ async def test_get_team_membership_waiter_cancel_does_not_cancel_shared_load():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=_slow_find_unique)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
cache = UserApiKeyCache()
|
||||
|
||||
async def _load():
|
||||
|
|
@ -7468,6 +7536,7 @@ async def test_get_team_membership_waiter_cancel_does_not_cancel_shared_load():
|
|||
assert result is not None
|
||||
assert result.user_id == "u-shield"
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique.assert_awaited_once()
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -7486,6 +7555,7 @@ async def test_invalidate_team_member_spend_state_evicts_the_negative_cache_sent
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(side_effect=[None, membership_row])
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
before = await get_team_membership(
|
||||
user_id="u-1", team_id="t-1", prisma_client=mock_prisma_client, user_api_key_cache=cache
|
||||
|
|
@ -7516,6 +7586,7 @@ async def test_get_access_object_db_fetch_returns_validated_access_group():
|
|||
}
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
mock_prisma_client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=access_row)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
|
|
@ -7544,6 +7615,7 @@ async def test_get_team_object_by_alias_db_fetch_returns_cached_obj():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -7573,6 +7645,7 @@ async def test_get_team_object_by_alias_loads_model_aliases_relation():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -7604,6 +7677,7 @@ async def test_get_org_object_by_alias_db_fetch_returns_validated_org():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[org_row])
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -7629,6 +7703,7 @@ async def test_get_object_permission_db_fetch_returns_validated_permission():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=perm_row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -7655,6 +7730,7 @@ async def test_get_managed_vector_store_rows_by_uuids_db_fetch_validates_rows():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[vs_row])
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -7682,6 +7758,7 @@ async def test_get_project_object_db_fetch_returns_cached_obj():
|
|||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=project_row)
|
||||
mock_prisma_client.replica_db = mock_prisma_client.db
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -8799,6 +8876,8 @@ class _MissingUserPrisma:
|
|||
async def find_unique(where: dict[str, str], include: dict[str, bool]) -> None:
|
||||
return None
|
||||
|
||||
replica_db = db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enforced_model_allowlists_treats_a_missing_user_row_as_unrestricted():
|
||||
|
|
@ -8824,6 +8903,8 @@ class _UnreachableUserPrisma:
|
|||
async def find_unique(where: dict[str, str], include: dict[str, bool]) -> None:
|
||||
raise RuntimeError("database gone")
|
||||
|
||||
replica_db = db
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enforced_model_allowlists_surfaces_a_failed_user_lookup():
|
||||
|
|
@ -8932,6 +9013,7 @@ async def test_access_group_model_fallback_uses_the_injected_database(channel: s
|
|||
)
|
||||
reader: Final = AsyncMock(return_value=group)
|
||||
client: Final = MagicMock(db=MagicMock(litellm_accessgrouptable=MagicMock(find_unique=reader)))
|
||||
client.replica_db = client.db
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", None), # test-quality-ok: [TQ008] prove reads stay on the injected connection
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", UserApiKeyCache()), # test-quality-ok: [TQ008] isolate the process cache
|
||||
|
|
@ -9076,6 +9158,7 @@ async def test_team_member_budget_check_temp_budget_increase_extends_cap():
|
|||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
async def mock_get_current_spend(counter_key, fallback_spend, max_budget=None, **kwargs):
|
||||
if counter_key == "spend:team_member:test-user:test-team":
|
||||
|
|
|
|||
|
|
@ -235,6 +235,7 @@ class FakeDb:
|
|||
class FakePrismaClient:
|
||||
def __init__(self, objects: dict[str, StoredObject]) -> None:
|
||||
self.db = FakeDb(objects)
|
||||
self.replica_db = self.db
|
||||
|
||||
|
||||
class FakeRouter:
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ def _wire_team_create_tx(prisma_client):
|
|||
)
|
||||
|
||||
prisma_client.db.tx = lambda *_args, **_kwargs: _tx()
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
|
||||
def test_microsoft_sso_handler_openid_from_response_user_principal_name():
|
||||
|
|
@ -595,6 +596,7 @@ async def test_default_team_params(team_params):
|
|||
# Mock Prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_teamtable.create = AsyncMock()
|
||||
_wire_team_create_tx(mock_prisma)
|
||||
mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0)
|
||||
|
|
@ -642,6 +644,7 @@ async def test_default_team_params_organization_id_reaches_sso_created_team(team
|
|||
litellm.default_team_params = team_params
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_teamtable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_teamtable.create = AsyncMock()
|
||||
_wire_team_create_tx(mock_prisma)
|
||||
|
|
@ -691,6 +694,7 @@ async def test_create_team_without_default_params():
|
|||
# Mock Prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_teamtable.find_first = AsyncMock(return_value=None)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_teamtable.create = AsyncMock()
|
||||
_wire_team_create_tx(mock_prisma)
|
||||
mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0)
|
||||
|
|
@ -973,6 +977,7 @@ async def test_upsert_sso_user_updates_role_for_existing_user():
|
|||
|
||||
# Mock prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
|
||||
|
||||
# Existing user in DB with old role
|
||||
|
|
@ -1023,6 +1028,7 @@ async def test_upsert_sso_user_does_not_update_invalid_role():
|
|||
|
||||
# Mock prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
|
||||
|
||||
# Existing user in DB
|
||||
|
|
@ -1068,6 +1074,7 @@ async def test_upsert_sso_user_no_role_in_sso_response():
|
|||
|
||||
# Mock prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
|
||||
|
||||
# Existing user in DB
|
||||
|
|
@ -1427,6 +1434,7 @@ async def test_check_and_update_if_proxy_admin_id():
|
|||
# Mock Prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_usertable.update = AsyncMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
||||
# Set up test data
|
||||
test_user_id = "test_admin_123"
|
||||
|
|
@ -1459,6 +1467,7 @@ async def test_check_and_update_if_proxy_admin_id_already_admin():
|
|||
# Mock Prisma client
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_usertable.update = AsyncMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
||||
# Set up test data
|
||||
test_user_id = "test_admin_123"
|
||||
|
|
@ -3196,6 +3205,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
for team_id in ("team1", "team2")
|
||||
]
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
with (
|
||||
patch.dict(
|
||||
os.environ,
|
||||
|
|
@ -3602,6 +3612,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
find_many = AsyncMock(return_value=[team_row])
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_many = find_many
|
||||
prisma_client.replica_db = prisma_client.db
|
||||
|
||||
details = await fetch_cli_sso_team_details(
|
||||
prisma_client=prisma_client, teams=["team-a"]
|
||||
|
|
@ -3633,6 +3644,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
failing_client.db.litellm_teamtable.find_many = AsyncMock(
|
||||
side_effect=Exception("connection reset")
|
||||
)
|
||||
failing_client.replica_db = failing_client.db
|
||||
assert (
|
||||
await fetch_cli_sso_team_details(
|
||||
prisma_client=failing_client, teams=["team-a"]
|
||||
|
|
@ -3642,6 +3654,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
empty_client = MagicMock()
|
||||
empty_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
|
||||
empty_client.replica_db = empty_client.db
|
||||
assert (
|
||||
await fetch_cli_sso_team_details(
|
||||
prisma_client=empty_client, teams=["team-a"]
|
||||
|
|
@ -6556,6 +6569,7 @@ async def test_setup_team_mappings():
|
|||
"""Test _setup_team_mappings function loads team mappings from database."""
|
||||
# Arrange
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_sso_config = MagicMock()
|
||||
mock_sso_config.sso_settings = {"team_mappings": {"team_ids_jwt_field": "groups"}}
|
||||
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(
|
||||
|
|
@ -7245,6 +7259,7 @@ class TestCliSsoAttributionMetadata:
|
|||
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=MagicMock(metadata={"auth_provider": "generic"})
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_usertable.update_many = AsyncMock()
|
||||
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(
|
||||
return_value=[
|
||||
|
|
@ -7538,6 +7553,7 @@ class TestSyncUserRoleFromJwtRoleMap:
|
|||
cache = DualCache()
|
||||
prisma = AsyncMock()
|
||||
prisma.db.litellm_usertable.update = AsyncMock()
|
||||
prisma.replica_db = prisma.db
|
||||
user_id = "testuser@example.com"
|
||||
|
||||
existing_user = LiteLLM_UserTable(
|
||||
|
|
@ -7577,6 +7593,7 @@ class TestSyncUserRoleFromJwtRoleMap:
|
|||
handler = self._make_jwt_handler()
|
||||
prisma = AsyncMock()
|
||||
prisma.db.litellm_usertable.update = AsyncMock()
|
||||
prisma.replica_db = prisma.db
|
||||
|
||||
existing_user = LiteLLM_UserTable(
|
||||
user_id="testuser@example.com",
|
||||
|
|
|
|||
|
|
@ -47,6 +47,7 @@ def _make_prisma_with_config(
|
|||
|
||||
client = MagicMock()
|
||||
client.db = db
|
||||
client.replica_db = db
|
||||
return client
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ from .conftest import normalize # type: ignore[import-not-found]
|
|||
def prisma_with_query_raw(monkeypatch):
|
||||
pc = MagicMock()
|
||||
pc.db.query_raw = AsyncMock(return_value=[])
|
||||
pc.replica_db = pc.db
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", pc)
|
||||
return pc
|
||||
|
||||
|
|
@ -195,6 +196,7 @@ def _alerting_client(
|
|||
row = MagicMock()
|
||||
row.param_value = db_row
|
||||
pc.db.litellm_config.find_first = AsyncMock(return_value=row)
|
||||
pc.replica_db = pc.db
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", pc)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
|
|
@ -283,6 +285,7 @@ def test_alerting_settings_reports_config_source_when_db_disagrees(
|
|||
row = MagicMock()
|
||||
row.param_value = {"alerting_args": db_alerting_args}
|
||||
pc.db.litellm_config.find_first = AsyncMock(return_value=row)
|
||||
pc.replica_db = pc.db
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", pc)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
|
|
@ -320,6 +323,7 @@ def test_alerting_settings_handles_empty_db_args(
|
|||
row = MagicMock()
|
||||
row.param_value = {"alerting_args": db_alerting_args}
|
||||
pc.db.litellm_config.find_first = AsyncMock(return_value=row)
|
||||
pc.replica_db = pc.db
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", pc)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
|
|
@ -376,6 +380,7 @@ def test_alerting_settings_happy(client, auth_as, monkeypatch):
|
|||
"""Pins ``GET /alerting/settings`` (happy: returns list of ConfigList entries)."""
|
||||
pc = MagicMock()
|
||||
pc.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||||
pc.replica_db = pc.db
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", pc)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
|
|
|
|||
|
|
@ -170,6 +170,7 @@ def _prisma_with_models(rows, existing_sentinel_rows=()):
|
|||
daily.upsert = AsyncMock()
|
||||
daily.delete_many = AsyncMock()
|
||||
prisma.db = types.SimpleNamespace(litellm_proxymodeltable=model_table, litellm_dailyteamspend=daily)
|
||||
prisma.replica_db = prisma.db
|
||||
return prisma, daily
|
||||
|
||||
|
||||
|
|
@ -757,6 +758,7 @@ def _prisma_for(model_rows, daily_table):
|
|||
model_table = MagicMock()
|
||||
model_table.find_many = AsyncMock(return_value=model_rows)
|
||||
prisma.db = types.SimpleNamespace(litellm_proxymodeltable=model_table, litellm_dailyteamspend=daily_table)
|
||||
prisma.replica_db = prisma.db
|
||||
return prisma
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -17376,11 +17376,12 @@ class TestMemberAutoRouterInference:
|
|||
user_id="router-member", team_id="router-team", user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
models=["member-router", "permitted-model"], api_key="test-key-hash", config={"timeout": 60},
|
||||
)
|
||||
self.database = SimpleNamespace(db=SimpleNamespace(
|
||||
tables = SimpleNamespace(
|
||||
litellm_teamtable=SimpleNamespace(find_unique=AsyncMock(return_value=self.team)),
|
||||
litellm_teammembership=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
|
||||
litellm_accessgrouptable=SimpleNamespace(find_unique=AsyncMock()),
|
||||
))
|
||||
)
|
||||
self.database = SimpleNamespace(db=tables, replica_db=tables)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", self.cache)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", self.database)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue