diff --git a/tests/integration/mcp/test_mcp_user_env_vars.py b/tests/integration/mcp/test_mcp_user_env_vars.py index d9cecaccadb..0170eb77491 100644 --- a/tests/integration/mcp/test_mcp_user_env_vars.py +++ b/tests/integration/mcp/test_mcp_user_env_vars.py @@ -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) diff --git a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py index d3e668b8987..7d13debd0cf 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py @@ -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() diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py index 6ddb8cbaa7c..8536fea60bc 100644 --- a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py +++ b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py @@ -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) diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py b/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py index 7f930f90247..e9146c84639 100644 --- a/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py +++ b/tests/test_litellm/integrations/cloudzero/test_cloudzero_database.py @@ -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 diff --git a/tests/test_litellm/integrations/focus/test_focus_database.py b/tests/test_litellm/integrations/focus/test_focus_database.py index 06240eac387..1aba9a9f90e 100644 --- a/tests/test_litellm/integrations/focus/test_focus_database.py +++ b/tests/test_litellm/integrations/focus/test_focus_database.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index cfcff73b857..0f5f60df934 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -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() diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index aa4bb2e49d3..f6c7faaf28d 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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": diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index 6f2341a578c..7913cc14f7f 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -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: diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 1230c548281..5b8cdf29b02 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -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", diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py b/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py index 7ef29b71bf0..2eecd2e7966 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_anthropic_beta.py @@ -47,6 +47,7 @@ def _make_prisma_with_config( client = MagicMock() client.db = db + client.replica_db = db return client diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py index b5536b7618c..cff6e61c1f5 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_metrics.py @@ -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() diff --git a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py index 8f25cffecf5..090e1af40cb 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py @@ -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 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 411377f29cf..d559171da23 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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)