diff --git a/tests/proxy_unit_tests/test_audit_logs_proxy.py b/tests/proxy_unit_tests/test_audit_logs_proxy.py index 878e19f5b6f..df0bdaa7248 100644 --- a/tests/proxy_unit_tests/test_audit_logs_proxy.py +++ b/tests/proxy_unit_tests/test_audit_logs_proxy.py @@ -178,6 +178,7 @@ async def test_create_audit_log_for_update_premium_user(): ): mock_prisma.db.litellm_auditlog.create = AsyncMock() + mock_prisma.replica_db = mock_prisma.db request_data = LiteLLM_AuditLogs( id="test_id", diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 2538556d3b5..96a203a83b7 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -826,7 +826,9 @@ async def test_get_fuzzy_user_object(): # Setup mock Prisma client mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_usertable = MagicMock() # Mock user data diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 6417c7c8aa6..d87fa61ed44 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -91,6 +91,7 @@ class TestCheckBatchCost: def mock_prisma_client(self): client = MagicMock() client.db = MagicMock() + client.replica_db = client.db client.db.litellm_managedobjecttable = MagicMock() client.db.litellm_usertable = MagicMock() return client @@ -1997,6 +1998,7 @@ class TestUnmanagedVertexRouting: prisma = instance.prisma_client prisma.db = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_managedobjecttable = MagicMock() prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) prisma.db.litellm_managedobjecttable.update = AsyncMock() @@ -2227,6 +2229,7 @@ class TestUnmanagedBedrockRouting: prisma = instance.prisma_client prisma.db = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_managedobjecttable = MagicMock() prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) prisma.db.litellm_managedobjecttable.update = AsyncMock() @@ -2410,6 +2413,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: proxy_logging_obj.get_proxy_hook.return_value = hook prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) instance = CheckBatchCost( @@ -2523,6 +2527,7 @@ class TestBatchCostAttribution: from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_row) prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) @@ -2807,6 +2812,7 @@ class TestPollPageStarvation: def _prisma(self, jobs): prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0) prisma.db.litellm_managedobjecttable.update = AsyncMock() prisma.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=jobs) @@ -3119,6 +3125,7 @@ class TestMultiPodBatchCostClaim: @staticmethod def _prisma(row: _FakeManagedObjectRow, journal: list): prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_managedobjecttable = _FakeManagedObjectTable(row, journal) prisma.db.litellm_managedfiletable.find_many = AsyncMock(return_value=[]) prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index e806e9a3394..1d335bd1cac 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -20,6 +20,7 @@ class TestCheckResponsesCost: """Create a mock Prisma client""" client = MagicMock() client.db = MagicMock() + client.replica_db = client.db client.db.litellm_managedobjecttable = MagicMock() return client diff --git a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py index edd0409343a..02aeb5792cd 100644 --- a/tests/proxy_unit_tests/test_default_end_user_budget_simple.py +++ b/tests/proxy_unit_tests/test_default_end_user_budget_simple.py @@ -46,6 +46,7 @@ async def test_default_budget_applied_to_end_user_without_budget(): } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) ) @@ -104,6 +105,7 @@ async def test_explicit_budget_not_overridden_by_default(): } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) ) @@ -161,6 +163,7 @@ async def test_budget_enforcement_blocks_over_budget_users(): } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) ) @@ -219,6 +222,7 @@ async def test_system_works_without_default_budget_configured(): } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock( return_value=MagicMock(dict=lambda: mock_end_user_data) ) diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index e95ed42013b..b9ad9dec831 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -47,6 +47,7 @@ async def test_jwt_to_virtual_key_mapping_resolution(): jwt_claims = {"email": "user@example.com", "sub": "123"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock() # Mock finding a mapping @@ -123,6 +124,7 @@ async def test_colliding_claim_value_from_another_issuer_does_not_resolve_to_the return None prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(side_effect=fake_find_first) # Dependency-inject the resolved key via the cache (IdentityStore._resolve_key @@ -206,6 +208,7 @@ async def test_global_mapping_resolution_is_cached_under_the_global_key_not_the_ return None prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db find_first = AsyncMock(side_effect=fake_find_first) prisma_client.db.litellm_jwtkeymapping.find_first = find_first @@ -251,6 +254,7 @@ async def test_jwt_to_virtual_key_mapping_no_mapping(): jwt_claims = {"email": "unknown@example.com"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock() prisma_client.db.litellm_jwtkeymapping.find_first.return_value = None @@ -413,6 +417,7 @@ def _make_non_admin_auth() -> UserAPIKeyAuth: def _mock_prisma(): prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_jwtkeymapping.create = AsyncMock() prisma.db.litellm_jwtkeymapping.find_unique = AsyncMock() prisma.db.litellm_jwtkeymapping.find_many = AsyncMock() @@ -688,6 +693,7 @@ async def test_reject_behavior_raises_403_on_no_mapping(): jwt_claims = {"email": "unknown@example.com"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) user_api_key_cache = DualCache() @@ -727,6 +733,7 @@ async def test_reject_behavior_caches_sentinel_after_db_miss(): jwt_claims = {"email": "unknown@example.com"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) user_api_key_cache = DualCache() @@ -784,6 +791,7 @@ async def test_reject_behavior_raises_403_on_cached_no_mapping(): jwt_claims = {"email": "unknown@example.com"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) # Pre-populate the negative cache so the DB is not hit @@ -831,6 +839,7 @@ async def test_auto_register_returns_pending_signal_without_creating_key(): jwt_claims = {"sub": "new-user-42"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() @@ -876,6 +885,7 @@ async def test_auto_register_creates_key_and_mapping_when_helper_invoked(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() @@ -946,6 +956,7 @@ async def test_auto_register_returns_pending_signal_on_stale_no_mapping_sentinel jwt_claims = {"email": "alice@corp.com"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock(return_value=None) prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() @@ -999,6 +1010,7 @@ async def test_auto_register_race_condition_unique_conflict(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.create = AsyncMock( side_effect=Exception("Unique constraint failed (P2002)") ) @@ -1284,6 +1296,7 @@ async def test_auto_register_race_conflict_tolerates_delete_failure(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.create = AsyncMock( side_effect=Exception("Unique constraint failed (P2002)") ) @@ -1350,6 +1363,7 @@ async def test_auto_register_raises_503_when_winner_mapping_vanishes(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.create = AsyncMock( side_effect=Exception("Unique constraint failed (P2002)") ) @@ -1404,6 +1418,7 @@ async def test_proxy_admin_sentinel_skips_db_lookup_on_cache_hit(): jwt_claims = {"sub": "admin-user"} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db # Will fail the test if accessed — proves the sentinel short-circuits DB prisma_client.db.litellm_jwtkeymapping.find_first = AsyncMock( side_effect=AssertionError("DB must not be hit when sentinel is cached") @@ -1451,6 +1466,7 @@ async def test_auto_register_helper_stamps_validated_identity_context(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() mock_key_obj = UserAPIKeyAuth( token="hashed", team_id="validated-team", user_id="validated-user" diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index ed0380058a5..a0fe65a15a1 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -2400,6 +2400,7 @@ async def test_proxy_server_prisma_setup(): mock_db = MagicMock() mock_db.start_token_refresh_task = AsyncMock() mock_client.db = mock_db + mock_client.replica_db = mock_client.db prisma_client = await ProxyStartupEvent._setup_prisma_client( database_url=os.getenv("DATABASE_URL"), @@ -3062,6 +3063,7 @@ async def test_update_config_success_callback_normalization(): class MockPrisma: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_config = MagicMock() self.db.litellm_config.find_first = AsyncMock(side_effect=fake_find_first) self.db.litellm_config.upsert = AsyncMock(side_effect=fake_upsert) diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 1134f41a940..4e6aadda02b 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -1642,6 +1642,7 @@ class MockPrismaClientDB: mock_key_data, ): self.db = MockDb(mock_team_data, mock_key_data) + self.replica_db = self.db async def get_data( self, @@ -1848,6 +1849,7 @@ async def test_health_check_not_called_when_disabled(monkeypatch): # Create mock prisma client mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.connect = AsyncMock() mock_prisma.health_check = AsyncMock() mock_prisma.check_view_exists = AsyncMock() @@ -1857,6 +1859,7 @@ async def test_health_check_not_called_when_disabled(monkeypatch): mock_db = MagicMock() mock_db.start_token_refresh_task = AsyncMock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock PrismaClient constructor monkeypatch.setattr( "litellm.proxy.proxy_server.PrismaClient", lambda **kwargs: mock_prisma @@ -2284,6 +2287,7 @@ def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch): def mock_prisma_client(): client = MagicMock() client.db = MagicMock() + client.replica_db = client.db client.db.litellm_teamtable = AsyncMock() return client diff --git a/tests/proxy_unit_tests/test_update_spend.py b/tests/proxy_unit_tests/test_update_spend.py index ebe505b3d60..8e5a3fb9057 100644 --- a/tests/proxy_unit_tests/test_update_spend.py +++ b/tests/proxy_unit_tests/test_update_spend.py @@ -29,6 +29,7 @@ class MockPrismaClient: def __init__(self): # Create AsyncMock for db operations self.db = AsyncMock() + self.replica_db = self.db self.db.litellm_spendlogs = AsyncMock() self.db.litellm_spendlogs.create_many = AsyncMock() diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index 9cdac341b1f..e4f274497be 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -1519,6 +1519,7 @@ async def test_user_budget_lookup_tolerates_an_unreadable_user(): from litellm.proxy.auth.user_api_key_auth import _read_user_model_max_budget prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db with patch( "litellm.proxy.auth.user_api_key_auth.get_user_object",