From 0ae887c1e78ff15bdeaae29f70a6de48e34344ef Mon Sep 17 00:00:00 2001 From: yuneng Date: Thu, 24 Sep 2026 11:58:01 +0000 Subject: [PATCH] test: expose replica_db on remaining PrismaClient test doubles Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../send_emails/test_base_email.py | 5 ++ .../send_emails/test_endpoints.py | 6 +++ .../test_batch_retrieve_input_file_id.py | 1 + ...trieve_registers_missing_output_file_id.py | 3 ++ ..._retrieve_returns_unified_input_file_id.py | 3 ++ .../test_deleted_file_returns_403_not_404.py | 3 ++ .../proxy/test_file_deletion_blocking.py | 1 + .../proxy/test_managed_files_access_check.py | 3 ++ .../auth/test_user_api_key_auth_mcp.py | 8 +++ .../mcp_server/test_byok_oauth_endpoints.py | 1 + .../mcp_server/test_mcp_env_vars.py | 16 ++++++ .../mcp_server/test_mcp_partial_update.py | 27 ++++++++++ .../mcp_server/test_mcp_server.py | 6 +++ .../mcp_server/test_mcp_sigv4_auth.py | 8 +++ .../mcp_server/test_proxy_api_credentials.py | 2 + .../agent_endpoints/test_agent_registry.py | 17 +++++++ .../test_analytics_endpoints.py | 2 + .../test_claude_code_marketplace.py | 1 + .../auth/test_admin_viewer_handler_access.py | 1 + .../test_auth_hot_path_network_requests.py | 5 ++ .../auth/test_custom_auth_end_user_budget.py | 4 ++ .../proxy/auth/test_login_utils.py | 22 +++++++++ .../auth/test_model_access_group_budgets.py | 2 + .../proxy/auth/test_onboarding.py | 13 +++++ .../proxy/batches_endpoints/test_endpoints.py | 1 + .../common_utils/test_config_sync_pubsub.py | 5 ++ ..._expired_ui_session_key_cleanup_manager.py | 1 + .../common_utils/test_key_rotation_e2e.py | 8 +++ .../test_key_rotation_integration.py | 2 + .../common_utils/test_key_rotation_manager.py | 5 ++ .../test_periodic_reload_schedule.py | 3 ++ .../test_registry_read_through.py | 14 ++++++ .../proxy/db/mcp_server/test_db.py | 3 ++ .../proxy/db/test_spend_counter_reseed.py | 4 +- .../proxy/db/test_spend_log_tool_index.py | 8 +++ .../proxy/db/test_tool_registry_writer.py | 12 +++++ .../guardrails/test_guardrail_endpoints.py | 30 ++++++++++++ .../guardrails/test_guardrail_registry.py | 1 + .../proxy/guardrails/test_usage_endpoints.py | 8 +++ .../proxy/guardrails/test_usage_tracking.py | 10 ++++ .../hooks/test_user_management_event_hooks.py | 1 + .../scim/test_scim_key_deactivation.py | 1 + .../scim/test_scim_patch_user.py | 4 ++ .../scim/test_scim_transformations.py | 1 + .../test_search_tool_management.py | 1 + .../test_access_group_endpoints.py | 7 +++ .../test_access_group_management.py | 12 +++++ .../test_activity_tenant_scoping.py | 7 +++ .../test_budget_endpoints.py | 1 + .../test_cache_settings_endpoints.py | 16 ++++++ .../management_endpoints/test_common_utils.py | 5 +- .../test_config_override_endpoints.py | 1 + .../test_coordination_redis_endpoints.py | 3 ++ .../test_customer_budget.py | 5 ++ .../test_customer_endpoints.py | 26 ++++++++++ .../test_delete_verification_tokens_failed.py | 2 + .../test_mcp_management_endpoints.py | 5 ++ .../test_organization_endpoints.py | 30 +++++++++++- .../test_password_endpoints.py | 9 ++++ .../test_project_org_authz.py | 8 +++ .../test_session_endpoints.py | 5 ++ .../test_tag_management_endpoints.py | 27 ++++++++-- .../test_team_callback_endpoints.py | 14 ++++++ .../test_team_default_params.py | 8 +++ .../test_team_model_alias_merge.py | 1 + .../test_tool_management_endpoints.py | 6 +++ .../test_workflow_management_endpoints.py | 34 +++++++++++++ .../test_audit_log_callbacks.py | 3 ++ .../test_bulk_user_creation.py | 4 ++ .../test_bulk_user_deletion.py | 2 + .../test_management_helpers_utils.py | 18 +++++++ .../test_object_permission_utils.py | 7 +++ .../test_team_metadata_validation.py | 2 + .../proxy/memory/test_memory_endpoints.py | 1 + .../test_files_common_utils.py | 7 +++ .../test_files_endpoint.py | 1 + .../test_managed_id_rewriter.py | 7 +++ .../test_policy_engine_endpoints.py | 7 +++ .../policy_engine/test_policy_validator.py | 1 + .../policy_engine/test_policy_versioning.py | 16 ++++++ .../test_policy_versioning_e2e.py | 2 + .../proxy/prompts/test_prompt_endpoints.py | 1 + .../prompts/test_prompt_endpoints_crud.py | 6 +++ .../proxy/prompts/test_prompt_environment.py | 3 ++ .../proxy/proxy_server/test_lifecycle.py | 3 ++ .../test_routes_model_cost_map.py | 1 + .../proxy/proxy_server/test_spend_counters.py | 5 ++ .../proxy/rag_endpoints/test_rag_endpoints.py | 4 ++ .../test_cloudzero_endpoints.py | 5 ++ .../test_spend_counter_batch.py | 4 ++ .../test_spend_query_optimization.py | 11 +++++ .../proxy/test_budget_reservation.py | 4 +- .../test_fallback_management_endpoints.py | 4 ++ ...test_filter_models_by_team_access_group.py | 4 ++ .../proxy/test_health_check_functions.py | 9 ++++ .../proxy/test_route_a2a_models.py | 1 + .../proxy/test_route_llm_request.py | 4 +- .../proxy/test_team_member_update.py | 1 + .../test_proxy_setting_endpoints.py | 49 +++++++++++++++++++ .../test_user_banner_endpoints.py | 4 ++ .../proxy/utils/prisma_and_spend/conftest.py | 1 + .../test_config_param_cache.py | 3 ++ .../prisma_and_spend/test_password_helpers.py | 3 ++ .../test_proxy_update_spend.py | 21 ++++++++ .../prisma_and_spend/test_spend_functions.py | 13 +++++ .../test_vector_store_access_control.py | 2 + .../test_vector_store_endpoints.py | 20 ++++++++ .../adaptive_router/test_adaptive_router.py | 3 ++ .../test_e2e_adaptive_router.py | 3 ++ .../adaptive_router/test_update_queue.py | 2 + .../test_litellm/test_model_block_unblock.py | 1 + .../test_router_retry_policy_update.py | 1 + .../test_vector_store_registry.py | 1 + 113 files changed, 781 insertions(+), 9 deletions(-) diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 8b89c592f02..d9fd4d71862 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -340,6 +340,7 @@ async def test_get_invitation_link(base_email_logger): return [mock_invitation_row] mock_prisma.db.litellm_invitationlink.find_many = mock_find_many + mock_prisma.replica_db = mock_prisma.db with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): # Test with valid user_id @@ -361,6 +362,7 @@ async def test_get_invitation_link(base_email_logger): return [] mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_empty + mock_prisma.replica_db = mock_prisma.db result = await base_email_logger._get_invitation_link( user_id="test-user", base_url="http://test.com" ) @@ -386,6 +388,7 @@ async def test_get_invitation_link_creates_new_when_none_exist(base_email_logger return [] mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_empty + mock_prisma.replica_db = mock_prisma.db # Mock the create_invitation_for_user function mock_created_invitation = mock.MagicMock() @@ -428,6 +431,7 @@ async def test_get_invitation_link_uses_existing_when_available(base_email_logge return [mock_invitation_row] mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_existing + mock_prisma.replica_db = mock_prisma.db with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with mock.patch( @@ -459,6 +463,7 @@ async def test_get_invitation_link_creates_new_when_list_is_none(base_email_logg return None mock_prisma.db.litellm_invitationlink.find_many = mock_find_many_none + mock_prisma.replica_db = mock_prisma.db # Mock the create_invitation_for_user function mock_created_invitation = mock.MagicMock() diff --git a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py index 7b32d9e8c44..a4ac4cc2dba 100644 --- a/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py +++ b/tests/test_litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py @@ -51,6 +51,7 @@ def mock_prisma_client(): mock_db.litellm_config = mock_config mock_client.db = mock_db + mock_client.replica_db = mock_client.db return mock_client @@ -65,6 +66,7 @@ async def test_get_email_settings_empty_db(mock_prisma_client): return None mock_prisma_client.db.litellm_config.find_unique = mock_find_unique + mock_prisma_client.replica_db = mock_prisma_client.db # Call the function result = await _get_email_settings(mock_prisma_client) @@ -91,6 +93,7 @@ async def test_get_email_settings_with_existing_settings(mock_prisma_client): return mock_entry mock_prisma_client.db.litellm_config.find_unique = mock_find_unique + mock_prisma_client.replica_db = mock_prisma_client.db # Call the function result = await _get_email_settings(mock_prisma_client) @@ -109,12 +112,14 @@ async def test_save_email_settings_new_entry(mock_prisma_client): return None mock_prisma_client.db.litellm_config.find_unique = mock_find_unique + mock_prisma_client.replica_db = mock_prisma_client.db # Setup mock upsert to return None async def mock_upsert(*args, **kwargs): return None mock_prisma_client.db.litellm_config.upsert = mock_upsert + mock_prisma_client.replica_db = mock_prisma_client.db # Settings to save settings = { @@ -273,6 +278,7 @@ def _prisma_recording_upserts(upserts): return None client.db.litellm_config.find_unique = find_unique + client.replica_db = client.db client.db.litellm_config.upsert = upsert return client diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py index 6e9c3c0354b..9eabc20a8df 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py @@ -56,6 +56,7 @@ async def test_should_resolve_raw_input_file_id_to_unified(): mock_prisma = MagicMock() mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=mock_db_object) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=mock_managed_file) from litellm.proxy.openai_files_endpoints.common_utils import get_batch_from_database diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py index 55a4d05946e..a8abd8300cc 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_registers_missing_output_file_id.py @@ -49,6 +49,7 @@ def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ=") def _build_prisma_mock(): mock = MagicMock() mock.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) + mock.replica_db = mock.db return mock @@ -134,6 +135,7 @@ async def test_get_batch_from_database_registers_missing_output_file_id(): prisma.db.litellm_managedobjecttable.find_first = AsyncMock( return_value=batch_db_record ) + prisma.replica_db = prisma.db prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) mock_managed_files = _build_managed_files_mock(unified_id=unified_output_file_id) @@ -188,6 +190,7 @@ async def test_registered_output_file_row_denies_cross_user_access(): raw_output_file_id = "file-raw-output" prisma = MagicMock() prisma.db.litellm_managedfiletable.upsert = AsyncMock() + prisma.replica_db = prisma.db prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py index 5d1b63a815e..7dd18a9eb71 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_returns_unified_input_file_id.py @@ -26,6 +26,7 @@ def _mock_prisma(batch_json: str, managed_file_record=None): prisma.db.litellm_managedobjecttable.find_first = AsyncMock( return_value=batch_db_record ) + prisma.replica_db = prisma.db prisma.db.litellm_managedfiletable.find_first = AsyncMock( return_value=managed_file_record @@ -81,6 +82,7 @@ async def test_should_resolve_raw_input_file_id_to_unified_id(): prisma.db.litellm_managedfiletable.find_first.assert_any_call( where={"flat_model_file_ids": {"has": raw_input_file_id}} ) + prisma.replica_db = prisma.db prisma.db.litellm_managedfiletable.find_first.assert_any_call( where={"flat_model_file_ids": {"has": "file-output-raw"}} ) @@ -125,3 +127,4 @@ async def test_should_preserve_already_managed_input_file_id(): ) prisma.db.litellm_managedfiletable.find_first.assert_not_called() + prisma.replica_db = prisma.db diff --git a/tests/test_litellm/enterprise/proxy/test_deleted_file_returns_403_not_404.py b/tests/test_litellm/enterprise/proxy/test_deleted_file_returns_403_not_404.py index 7ad564dc8f9..0a0eebf7d6e 100644 --- a/tests/test_litellm/enterprise/proxy/test_deleted_file_returns_403_not_404.py +++ b/tests/test_litellm/enterprise/proxy/test_deleted_file_returns_403_not_404.py @@ -37,6 +37,7 @@ def _make_managed_files_with_no_db_record(): mock_prisma = MagicMock() mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db return _PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), @@ -76,6 +77,7 @@ async def test_should_allow_owner_access_when_record_exists(): mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), @@ -105,6 +107,7 @@ async def test_should_block_different_user_when_record_exists(): mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db managed_files = _PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), diff --git a/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py b/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py index 852077dcf0c..6e5235965ac 100644 --- a/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py +++ b/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py @@ -84,6 +84,7 @@ def _make_managed_files_instance_with_batches( mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock( return_value=mock_file_record ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_managedfiletable.delete = AsyncMock( return_value=mock_file_record ) diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py b/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py index ad46798b788..0370b8df4bf 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py @@ -52,6 +52,7 @@ def _make_managed_files_instance( mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db instance = _PROXY_LiteLLMManagedFiles( internal_usage_cache=MagicMock(), @@ -189,6 +190,7 @@ def _make_managed_files_instance_with_object_store(): mock_prisma = MagicMock() mock_prisma.db.litellm_managedobjecttable.upsert = AsyncMock(side_effect=upsert) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_managedobjecttable.find_first = AsyncMock( side_effect=find_first ) @@ -296,6 +298,7 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[mock_job] ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 05ed53df8e8..0ef2bd5c338 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -993,6 +993,7 @@ class TestMCPRequestHandler: # Test case: None values in database mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.find_unique.return_value = None + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique.return_value = None user_api_key_auth = UserAPIKeyAuth( @@ -4589,6 +4590,7 @@ class TestAgentMCPPermissions: agent_row.object_permission_id = "perm-xyz" prisma_client = MagicMock() prisma_client.db.litellm_agentstable.find_unique = AsyncMock(return_value=agent_row) + prisma_client.replica_db = prisma_client.db user_api_key_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", @@ -4627,6 +4629,7 @@ class TestAgentMCPPermissions: agent_row.object_permission_id = None prisma_client = MagicMock() prisma_client.db.litellm_agentstable.find_unique = AsyncMock(return_value=agent_row) + prisma_client.replica_db = prisma_client.db user_api_key_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", @@ -9421,6 +9424,7 @@ class TestGetUserObjectPermission: def _prisma_with_user(self, user_row): prisma_client = MagicMock() prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + prisma_client.replica_db = prisma_client.db return prisma_client async def test_resolves_through_the_shared_permission_cache(self): @@ -9447,6 +9451,7 @@ class TestGetUserObjectPermission: # The user_id -> object_permission_id link is cached, so the user row is read once. prisma_client.db.litellm_usertable.find_unique.reset_mock() + prisma_client.replica_db = prisma_client.db await MCPRequestHandler._get_user_object_permission(auth) prisma_client.db.litellm_usertable.find_unique.assert_not_called() @@ -9469,6 +9474,7 @@ class TestGetUserObjectPermission: assert await MCPRequestHandler._get_user_object_permission(auth) is None mock_get_perm.assert_not_awaited() prisma_client.db.litellm_usertable.find_unique.assert_awaited_once() + prisma_client.replica_db = prisma_client.db async def test_missing_user_row_places_no_ceiling(self): """Whether this human is entitled at all is unknown when their row is absent, which is the @@ -9490,6 +9496,7 @@ class TestGetUserObjectPermission: prisma_client = MagicMock() prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=Exception("db down")) + prisma_client.replica_db = prisma_client.db auth = UserAPIKeyAuth(api_key="sk-test", user_id="human-db-down") with ( @@ -9551,6 +9558,7 @@ def _agent_prisma(object_permission_id=None, side_effect=None): return_value=MagicMock(object_permission_id=object_permission_id), side_effect=side_effect, ) + prisma_client.replica_db = prisma_client.db return prisma_client diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index 6d2ea2ff301..ea1a819b276 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -633,6 +633,7 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True) prisma = MagicMock() prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None) + prisma.replica_db = prisma.db monkeypatch.setattr(proxy_server, "prisma_client", prisma) with pytest.raises(HTTPException) as exc_info: await mcp_operations.execute_mcp_tool( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py index 93b894f7645..8b53cbfd318 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -790,6 +790,7 @@ async def test_load_user_env_vars_force_refresh_bypasses_cache( prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock( side_effect=[old_row, new_row] ) + prisma.replica_db = prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) manager = MCPServerManager() @@ -836,6 +837,7 @@ async def test_load_user_env_vars_invalidation_forces_refetch( prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock( side_effect=[old_row, new_row] ) + prisma.replica_db = prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) manager = MCPServerManager() @@ -870,6 +872,7 @@ def _mock_env_vars_prisma(row=None): prisma = MagicMock() prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock(return_value=row) + prisma.replica_db = prisma.db prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[]) prisma.db.litellm_mcpuserenvvars.upsert = AsyncMock() prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock() @@ -965,6 +968,7 @@ def _transactional_env_vars_prisma(read_delay: float = 0.0): class _Prisma: def __init__(self, delay): self.db = _DB(_Store(), delay) + self.replica_db = self.db return _Prisma(read_delay) @@ -1057,6 +1061,7 @@ async def test_get_user_env_vars_bulk_distributes_results(env_vars_salt_key): prisma = _mock_env_vars_prisma() prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock(return_value=[row1, row2]) + prisma.replica_db = prisma.db result = await get_user_env_vars_bulk(prisma, "alice", ["srv-1", "srv-2", "srv-3"]) assert result == {"srv-1": {"A": "1"}, "srv-2": {"B": "2"}} @@ -1080,6 +1085,7 @@ async def test_delete_user_env_vars_is_idempotent_delete_many(): prisma = _mock_env_vars_prisma() await delete_user_env_vars(prisma, "alice", "srv-1") prisma.db.litellm_mcpuserenvvars.delete_many.assert_awaited_once() + prisma.replica_db = prisma.db call = prisma.db.litellm_mcpuserenvvars.delete_many.call_args assert call.kwargs["where"] == {"user_id": "alice", "server_id": "srv-1"} @@ -1187,6 +1193,7 @@ async def test_merge_user_env_vars_acquires_lock_without_deserializing_void( tx = _Tx() prisma = MagicMock() prisma.db.tx = MagicMock(return_value=tx) + prisma.replica_db = prisma.db values = {"CORP_TOKEN": "t0ken"} merged = await merge_user_env_vars( @@ -1207,6 +1214,7 @@ async def test_delete_mcp_server_removes_orphaned_user_env_vars(): prisma = _mock_env_vars_prisma() prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=object()) + prisma.replica_db = prisma.db await delete_mcp_server(prisma, "srv-1") @@ -1224,6 +1232,7 @@ async def test_delete_mcp_server_skips_env_var_cleanup_when_server_missing(): prisma = _mock_env_vars_prisma() prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=None) + prisma.replica_db = prisma.db result = await delete_mcp_server(prisma, "srv-1") @@ -1244,6 +1253,7 @@ async def test_delete_mcp_server_succeeds_when_orphan_cleanup_fails(): deleted = object() prisma = _mock_env_vars_prisma() prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=deleted) + prisma.replica_db = prisma.db prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock( side_effect=Exception("connection pool exhausted") ) @@ -1265,6 +1275,7 @@ async def test_delete_mcp_server_removes_orphaned_user_credentials(): prisma = _mock_env_vars_prisma() prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=object()) + prisma.replica_db = prisma.db await delete_mcp_server(prisma, "srv-1") @@ -1282,6 +1293,7 @@ async def test_delete_mcp_server_skips_credential_cleanup_when_server_missing(): prisma = _mock_env_vars_prisma() prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=None) + prisma.replica_db = prisma.db result = await delete_mcp_server(prisma, "srv-1") @@ -1301,6 +1313,7 @@ async def test_delete_mcp_server_credential_cleanup_failure_still_cleans_env_var deleted = object() prisma = _mock_env_vars_prisma() prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=deleted) + prisma.replica_db = prisma.db prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock( side_effect=Exception("connection pool exhausted") ) @@ -1534,6 +1547,7 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri mock_prisma.db.litellm_mcpservertable.create = AsyncMock( return_value=_prisma_row_with_json_string_env_vars() ) + mock_prisma.replica_db = mock_prisma.db created = await create_mcp_server( mock_prisma, @@ -1553,6 +1567,7 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri mock_prisma_upd.db.litellm_mcpservertable.update = AsyncMock( return_value=_prisma_row_with_json_string_env_vars() ) + mock_prisma_upd.replica_db = mock_prisma_upd.db updated = await update_mcp_server( mock_prisma_upd, UpdateMCPServerRequest(server_id="srv-update"), @@ -1630,6 +1645,7 @@ async def test_rotate_mcp_user_env_vars_logs_rotated_and_skipped_counts( prisma.db.litellm_mcpuserenvvars.find_many = AsyncMock( return_value=[good_one, good_two, bad] ) + prisma.replica_db = prisma.db prisma.db.litellm_mcpuserenvvars.update = AsyncMock() logger = MagicMock() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py index bd37a976286..a645f289d1d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -29,6 +29,7 @@ def _credentials_cleared(value) -> bool: def _mock_prisma(): mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable = AsyncMock() + mock_prisma.replica_db = mock_prisma.db row = models.LiteLLM_MCPServerTable.model_construct(server_id="test-server", transport="http", env={}, env_vars=[]) mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=row) mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=row) @@ -198,6 +199,7 @@ async def _run_update_with_existing(data: UpdateMCPServerRequest, existing_auth_ existing.auth_type = existing_auth_type existing.credentials = None mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db await update_mcp_server(mock_prisma, data, "test-user") return mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] @@ -240,6 +242,7 @@ async def test_explicit_null_clears_upstream_resource_and_keeps_the_rest_of_the_ existing.url = "https://up.example.com/mcp" existing.credentials = json.dumps({"client_secret": "csec", "upstream_resource": "api://audience"}) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="my-test-server", credentials={"upstream_resource": None}) await update_mcp_server(mock_prisma, data, "test-user") @@ -261,6 +264,7 @@ async def test_url_change_clears_stale_oauth_fields(): existing.url = "https://old.example.com/mcp" existing.credentials = None mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="my-test-server", url="https://new.example.com/mcp") await update_mcp_server(mock_prisma, data, "test-user") @@ -286,6 +290,7 @@ async def test_url_change_clears_stale_oauth_fields_even_when_resubmitted_unchan existing.token_url = "https://old-idp.example.com/token" existing.authorization_url = "https://old-idp.example.com/authorize" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="my-test-server", @@ -318,6 +323,7 @@ async def test_clearing_pinned_issuer_clears_stale_oauth_endpoints(): existing.token_url = "https://pinned-idp.example.com/token" existing.authorization_url = "https://pinned-idp.example.com/authorize" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="my-test-server", @@ -345,6 +351,7 @@ async def test_repointing_pinned_issuer_clears_stale_endpoints_keeps_new_issuer( existing.token_url = "https://old-idp.example.com/token" existing.authorization_url = "https://old-idp.example.com/authorize" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="my-test-server", @@ -375,6 +382,7 @@ async def test_establishing_issuer_first_time_preserves_endpoints_set_in_the_sam existing.credentials = None existing.issuer = None mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="my-test-server", @@ -402,6 +410,7 @@ async def test_unchanged_url_does_not_clear_oauth_fields(): existing.url = "https://same.example.com/mcp" existing.credentials = None mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="my-test-server", url="https://same.example.com/mcp") await update_mcp_server(mock_prisma, data, "test-user") @@ -600,6 +609,7 @@ async def test_credentials_merge_migrates_legacy_blob_te_settings(): }, ) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="te-server", @@ -629,6 +639,7 @@ async def test_cleared_column_is_not_resurrected_by_legacy_blob_value(): credentials={"client_id": "enc-old-cid", "token_exchange_endpoint": "https://dead-idp.example.com/token"}, ) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="te-server", @@ -654,6 +665,7 @@ async def test_merge_strips_blob_te_copy_when_column_already_set(): ) existing.token_exchange_endpoint = "https://column.example.com/token" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="te-server", @@ -676,6 +688,7 @@ async def test_auth_type_switch_clears_flow_fields_with_external_fields_set(): mock_prisma = _mock_prisma() existing = _existing_row("oauth2") mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="te-server", auth_type="oauth2_token_exchange") await update_mcp_server(mock_prisma, data, "test-user", fields_set=set(data.fields_set())) @@ -705,6 +718,7 @@ async def test_explicit_clear_without_credentials_purges_legacy_blob_copy(): credentials={"client_id": "enc-old-cid", "token_exchange_endpoint": "https://dead-idp.example.com/token"}, ) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="te-server", token_exchange_endpoint=None) await update_mcp_server(mock_prisma, data, "test-user") @@ -732,6 +746,7 @@ async def test_explicit_te_write_without_credentials_migrates_other_legacy_field }, ) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="te-server", audience="api://new") await update_mcp_server(mock_prisma, data, "test-user") @@ -751,6 +766,7 @@ async def test_te_update_without_blob_te_keys_leaves_credentials_untouched(): mock_prisma = _mock_prisma() existing = _existing_row("oauth2_token_exchange", credentials={"client_id": "enc-old-cid"}) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="te-server", token_exchange_endpoint="https://new.example.com/token") await update_mcp_server(mock_prisma, data, "test-user") @@ -774,6 +790,7 @@ async def test_cf_pair_switch_without_credentials_keeps_stored_app_and_endpoints existing.token_url = "https://provider.example/token" existing.registration_url = "https://provider.example/register" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="cf-server", auth_type="oauth_delegate") await update_mcp_server(mock_prisma, data, "test-user") @@ -791,6 +808,7 @@ async def test_cf_pair_switch_with_partial_credentials_merges_not_replaces(): mock_prisma = _mock_prisma() existing = _existing_row("true_passthrough", credentials={"client_id": "enc-A", "client_secret": "enc-B"}) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="cf-server", auth_type="oauth_delegate", credentials={"client_id": "B"}) await update_mcp_server(mock_prisma, data, "test-user") @@ -808,6 +826,7 @@ async def test_null_existing_auth_type_to_cf_counts_as_changed_and_clears_blob() mock_prisma = _mock_prisma() existing = _existing_row(None, credentials={"client_id": "enc-old"}) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest(server_id="cf-server", auth_type="true_passthrough") await update_mcp_server(mock_prisma, data, "test-user") @@ -827,6 +846,7 @@ async def test_client_rotation_strips_legacy_minted_token_keys(): "oauth2", credentials={"client_id": "A", "access_token": "T", "refresh_token": "R", "expires_in": 3600} ) mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="oauth2-server", auth_type="oauth2", credentials={"client_id": "B", "client_secret": "S"} @@ -872,6 +892,7 @@ def _mock_toolset_prisma(): } mock_prisma = MagicMock() mock_prisma.db.litellm_mcptoolsettable = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcptoolsettable.update = AsyncMock(return_value=updated_row) return mock_prisma @@ -945,6 +966,7 @@ async def test_find_identifier_conflict_reports_alias_hit(): mock_prisma = _mock_prisma() mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) + mock_prisma.replica_db = mock_prisma.db conflict = await find_mcp_server_identifier_conflict( mock_prisma, server_name="new-name", alias="taken", exclude_server_id="my-server" @@ -966,6 +988,7 @@ async def test_find_identifier_conflict_reports_server_name_when_alias_is_free() mock_prisma = _mock_prisma() mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(side_effect=[None, _conflict_row()]) + mock_prisma.replica_db = mock_prisma.db conflict = await find_mcp_server_identifier_conflict( mock_prisma, server_name="taken", alias="free", exclude_server_id=None @@ -994,6 +1017,7 @@ async def test_update_writing_alias_returns_conflict_instead_of_row(): mock_prisma = _mock_prisma() mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) + mock_prisma.replica_db = mock_prisma.db result = await update_mcp_server( mock_prisma, @@ -1040,6 +1064,7 @@ async def test_clearing_alias_conflicts_on_the_fallback_server_name(): existing = MagicMock() existing.server_name = "taken" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) result = await update_mcp_server( @@ -1063,6 +1088,7 @@ async def test_clearing_alias_to_empty_string_conflicts_on_the_fallback_server_n existing = MagicMock() existing.server_name = "taken" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=_conflict_row()) result = await update_mcp_server( @@ -1082,6 +1108,7 @@ async def test_clearing_alias_with_free_server_name_returns_the_row(): existing = MagicMock() existing.server_name = "free-name" mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db result = await update_mcp_server( mock_prisma, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ea82c9b069d..cb3b0001797 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5390,6 +5390,7 @@ class TestMCPServerManagerReload: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row]) + mock_prisma.replica_db = mock_prisma.db with ( patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", @@ -5432,6 +5433,7 @@ class TestMCPServerManagerReload: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row]) + mock_prisma.replica_db = mock_prisma.db with ( patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", @@ -5487,6 +5489,7 @@ class TestMCPServerManagerReload: mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock( return_value=[healthy_row, bad_row, another_healthy_row] ) + mock_prisma.replica_db = mock_prisma.db with ( patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", @@ -5557,6 +5560,7 @@ class TestMCPServerManagerReload: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[healthy_row, bad_openapi_row]) + mock_prisma.replica_db = mock_prisma.db with ( patch( "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", @@ -8718,6 +8722,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_queries_active_rows( row.server_id = "submitted-1" prisma_client = MagicMock() prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) + prisma_client.replica_db = prisma_client.db result = await get_active_submitted_mcp_server_ids_for_user(prisma_client, "submitter-user") @@ -8738,6 +8743,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_ prisma_client = MagicMock() prisma_client.db.litellm_mcpservertable.find_many = AsyncMock() + prisma_client.replica_db = prisma_client.db assert await get_active_submitted_mcp_server_ids_for_user(prisma_client, "") == [] prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index 8cf3bc6fcc7..0e2f48488b2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -605,6 +605,7 @@ class TestCredentialMergeOnUpdate: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row()) data = UpdateMCPServerRequest( @@ -645,6 +646,7 @@ class TestCredentialMergeOnUpdate: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row()) + mock_prisma.replica_db = mock_prisma.db data = UpdateMCPServerRequest( server_id="test-server", @@ -672,6 +674,7 @@ class TestCredentialMergeOnUpdate: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row()) data = UpdateMCPServerRequest( @@ -714,6 +717,7 @@ class TestCredentialMergeOnUpdate: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row()) data = UpdateMCPServerRequest( @@ -757,6 +761,7 @@ class TestCredentialMergeOnUpdate: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row()) data = UpdateMCPServerRequest( @@ -993,6 +998,7 @@ class TestRotateCredentials: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.update = AsyncMock() mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[]) @@ -1041,6 +1047,7 @@ class TestRotateCredentials: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.update = AsyncMock() mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[]) @@ -1088,6 +1095,7 @@ class TestAuthTypeSwitchClearsCredentials: mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row()) data = UpdateMCPServerRequest( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py index ed3e5f48516..632adc0fe25 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py @@ -142,6 +142,7 @@ async def test_mint_reads_the_users_teams_from_the_database_not_a_stale_cached_r prisma.db.litellm_usertable.find_unique = AsyncMock( return_value=_user(user_id="stale-cache-user", teams=["team-a"]) ) + prisma.replica_db = prisma.db monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) monkeypatch.setattr(proxy_server, "prisma_client", prisma) @@ -167,6 +168,7 @@ async def test_mint_refuses_a_user_scim_deactivated_after_the_cache_last_saw_the prisma.db.litellm_usertable.find_unique = AsyncMock( return_value=_user(user_id="deactivated-user", teams=["team-a"], metadata={"scim_active": False}) ) + prisma.replica_db = prisma.db monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) monkeypatch.setattr(proxy_server, "prisma_client", prisma) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py index b036e0dac4d..a58ab44c203 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -61,6 +61,7 @@ async def test_update_agent_in_db_clears_static_headers_and_extra_headers_when_o mock_update = AsyncMock(return_value=updated_agent) mock_prisma.db.litellm_agentstable.update = mock_update + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None) # Agent config WITHOUT static_headers or extra_headers (omitted) @@ -108,6 +109,7 @@ async def test_update_agent_in_db_preserves_explicit_static_headers_and_extra_he mock_update = AsyncMock(return_value=updated_agent) mock_prisma.db.litellm_agentstable.update = mock_update + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None) agent_config = { @@ -453,6 +455,7 @@ async def test_update_agent_in_db_raises_when_row_deleted_mid_update(): mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value=SimpleNamespace(litellm_params={}, object_permission_id=None) ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None) with pytest.raises(Exception, match="Error updating agent in DB") as exc_info: @@ -478,6 +481,7 @@ async def test_patch_agent_in_db_raises_when_row_deleted_mid_update(): mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value={"agent_id": "agent-123", "agent_name": "Old Agent", "object_permission_id": None} ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None) with pytest.raises(Exception, match="Error patching agent in DB") as exc_info: @@ -497,6 +501,7 @@ async def test_delete_agent_from_db_raises_when_row_already_gone(): registry: Final = AgentRegistry() mock_prisma: Final = MagicMock() mock_prisma.db.litellm_agentstable.delete = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db with pytest.raises(Exception, match="Error deleting agent from DB") as exc_info: await registry.delete_agent_from_db(agent_id="agent-123", prisma_client=mock_prisma) @@ -701,6 +706,7 @@ async def test_add_agent_to_db_drops_a_sentinel_value_instead_of_storing_the_pla created_agent.object_permission = None mock_create = AsyncMock(return_value=created_agent) mock_prisma.db.litellm_agentstable.create = mock_create + mock_prisma.replica_db = mock_prisma.db await registry.add_agent_to_db( agent={ @@ -738,6 +744,7 @@ async def test_update_agent_in_db_preserves_secret_when_echoed_back_redacted(): object_permission_id=None, ) ) + mock_prisma.replica_db = mock_prisma.db updated_agent = MagicMock() updated_agent.model_dump.return_value = { "agent_id": "agent-123", @@ -786,6 +793,7 @@ async def test_update_agent_in_db_preserves_secret_when_key_omitted_entirely(): object_permission_id=None, ) ) + mock_prisma.replica_db = mock_prisma.db updated_agent = MagicMock() updated_agent.model_dump.return_value = { "agent_id": "agent-123", @@ -832,6 +840,7 @@ async def test_update_agent_in_db_preserves_secret_nested_under_a_non_sensitive_ object_permission_id=None, ) ) + mock_prisma.replica_db = mock_prisma.db updated_agent = MagicMock() updated_agent.model_dump.return_value = { "agent_id": "agent-123", @@ -880,6 +889,7 @@ async def test_update_agent_in_db_clears_secret_on_explicit_empty_value(): object_permission_id=None, ) ) + mock_prisma.replica_db = mock_prisma.db updated_agent = MagicMock() updated_agent.model_dump.return_value = { "agent_id": "agent-123", @@ -922,6 +932,7 @@ async def test_patch_agent_in_db_preserves_secret_when_litellm_params_omitted(): "object_permission_id": None, } ) + mock_prisma.replica_db = mock_prisma.db patched_agent = MagicMock() patched_agent.model_dump.return_value = { "agent_id": "agent-123", @@ -964,6 +975,7 @@ async def test_patch_agent_in_db_preserves_secret_when_echoed_back_redacted(): "object_permission_id": None, } ) + mock_prisma.replica_db = mock_prisma.db patched_agent = MagicMock() patched_agent.model_dump.return_value = { "agent_id": "agent-123", @@ -1013,6 +1025,7 @@ async def test_add_agent_to_db_persists_deduplicated_access_group_ids(): mock_prisma: Final = MagicMock() mock_create = AsyncMock(return_value=_agent_row_mock(["ag-1", "ag-2"])) mock_prisma.db.litellm_agentstable.create = mock_create + mock_prisma.replica_db = mock_prisma.db result: Final = await registry.add_agent_to_db( agent={ @@ -1034,6 +1047,7 @@ async def test_add_agent_to_db_without_access_group_ids_leaves_column_to_its_def mock_prisma: Final = MagicMock() mock_create = AsyncMock(return_value=_agent_row_mock([])) mock_prisma.db.litellm_agentstable.create = mock_create + mock_prisma.replica_db = mock_prisma.db await registry.add_agent_to_db( agent={"agent_name": "Test Agent", "agent_card_params": _sample_agent_card_params()}, @@ -1067,6 +1081,7 @@ async def test_patch_agent_in_db_replaces_access_group_ids_when_provided( "access_group_ids": ["ag-1"], } ) + mock_prisma.replica_db = mock_prisma.db mock_update = AsyncMock(return_value=_agent_row_mock(expected)) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1090,6 +1105,7 @@ async def test_patch_agent_in_db_keeps_access_group_ids_when_omitted(): "access_group_ids": ["ag-1"], } ) + mock_prisma.replica_db = mock_prisma.db mock_update = AsyncMock(return_value=_agent_row_mock(["ag-1"])) mock_prisma.db.litellm_agentstable.update = mock_update @@ -1112,6 +1128,7 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, access_group_ids=["ag-1"]) ) + mock_prisma.replica_db = mock_prisma.db mock_update = AsyncMock(return_value=_agent_row_mock(expected)) mock_prisma.db.litellm_agentstable.update = mock_update body: Final = { diff --git a/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py index ea528878db8..d6ec3eb48d8 100644 --- a/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py +++ b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py @@ -53,6 +53,7 @@ MODEL_ROWS = [{"model": "gpt-5.1"}] def build_prisma(query_raw: AsyncMock) -> MagicMock: prisma = MagicMock() prisma.db.query_raw = query_raw + prisma.replica_db = prisma.db return prisma @@ -137,6 +138,7 @@ async def test_rejects_malformed_dates_with_400(mock_prisma: MagicMock): assert exc_info.value.status_code == 400 mock_prisma.db.query_raw.assert_not_called() + mock_prisma.replica_db = mock_prisma.db def test_totals_ratio_is_zero_without_requests(): diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py index a585666743f..21bdedaa076 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py @@ -75,6 +75,7 @@ def _make_mock_prisma(): mock_table.create = AsyncMock(side_effect=_create) mock_table.update = AsyncMock(side_effect=_update) mock_client.db.litellm_claudecodeplugintable = mock_table + mock_client.replica_db = mock_client.db return mock_client diff --git a/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py b/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py index 2309b0a931f..b2bbf2c1d6e 100644 --- a/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py +++ b/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py @@ -64,6 +64,7 @@ def admin_viewer_client(monkeypatch): litellm_config=mock_config_table, query_raw=mock_query_raw, ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) _override_auth(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) diff --git a/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py b/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py index e3b76cac8ce..224fa370cbe 100644 --- a/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py +++ b/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py @@ -254,6 +254,7 @@ async def test_get_team_object_warm_cache(): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable = MagicMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock() @@ -295,6 +296,7 @@ async def test_get_user_object_warm_cache(): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_usertable = MagicMock() mock_prisma.db.litellm_usertable.find_unique = AsyncMock() @@ -347,6 +349,7 @@ async def test_get_team_membership_warm_cache(): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teammembership = MagicMock() mock_prisma.db.litellm_teammembership.find_unique = AsyncMock() @@ -533,6 +536,7 @@ async def test_get_user_object_missing_user_negative_cache(): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_usertable = MagicMock() mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) @@ -564,6 +568,7 @@ async def test_get_user_object_missing_user_rechecks_after_expiry(): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_usertable = MagicMock() mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) diff --git a/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py index e83c5cf8419..ec9258b0a9b 100644 --- a/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py +++ b/tests/test_litellm/proxy/auth/test_custom_auth_end_user_budget.py @@ -201,6 +201,7 @@ async def test_custom_auth_token_budget_still_loads_and_caches_unrestricted_end_ 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_row) cache = UserApiKeyCache() @@ -241,6 +242,7 @@ async def test_custom_auth_key_default_end_user_budget_reaches_the_token_for_a_n 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=None) mock_prisma.db.litellm_budgettable.find_unique = AsyncMock(side_effect=_find_budget) @@ -279,6 +281,7 @@ async def test_custom_auth_cap_stays_below_the_key_default_end_user_budget(monke 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=None) mock_prisma.db.litellm_budgettable.find_unique = AsyncMock(side_effect=_find_budget) @@ -317,6 +320,7 @@ async def test_custom_auth_proxy_wide_default_end_user_budget_reaches_an_uncappe 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=None) mock_prisma.db.litellm_budgettable.find_unique = AsyncMock(side_effect=_find_budget) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 1b15994e777..4c4d6aa87f7 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -99,6 +99,7 @@ async def test_authenticate_user_admin_login_with_ui_credentials(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -150,6 +151,7 @@ async def test_authenticate_user_admin_login_with_master_key_as_password(monkeyp mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db env_vars = { "UI_USERNAME": ui_username, @@ -206,6 +208,7 @@ async def test_authenticate_user_invalid_credentials(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": "correct-password"}): with pytest.raises(ProxyException) as exc_info: @@ -260,6 +263,7 @@ async def test_authenticate_user_wrong_password(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=mock_user) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -311,6 +315,7 @@ async def test_authenticate_user_email_case_insensitive_login(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -366,6 +371,7 @@ async def test_authenticate_user_database_required_for_admin(monkeypatch): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password}): with patch( @@ -405,6 +411,7 @@ async def test_authenticate_user_admin_login_with_non_ascii_characters(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -479,6 +486,7 @@ async def test_authenticate_user_multiple_logins_generate_unique_tokens(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -566,6 +574,7 @@ async def test_authenticate_user_database_login_with_non_ascii_password(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(side_effect=mock_find_first) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -1768,6 +1777,7 @@ class TestDisablePasswordLoginWhenSSOEnabled: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": master_key}): with ExitStack() as stack: @@ -1801,6 +1811,7 @@ class TestDisablePasswordLoginWhenSSOEnabled: ) mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=mock_user) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict(os.environ, {"UI_USERNAME": "admin", "UI_PASSWORD": "unrelated"}): with ExitStack() as stack: @@ -1827,6 +1838,7 @@ class TestDisablePasswordLoginWhenSSOEnabled: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -1864,6 +1876,7 @@ class TestDisablePasswordLoginWhenSSOEnabled: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -1898,6 +1911,7 @@ class TestDisablePasswordLoginWhenSSOEnabled: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -1938,6 +1952,7 @@ class TestDisableEnvCredentialLogin: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict(os.environ, {"UI_USERNAME": ui_username, "UI_PASSWORD": ui_password}): with pytest.raises(ProxyException) as exc_info: @@ -1963,6 +1978,7 @@ class TestDisableEnvCredentialLogin: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict(os.environ, {"UI_USERNAME": "admin"}, clear=True): with pytest.raises(ProxyException) as exc_info: @@ -1991,6 +2007,7 @@ class TestDisableEnvCredentialLogin: ) mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=mock_user) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -2031,6 +2048,7 @@ class TestDisableEnvCredentialLogin: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch.dict( os.environ, @@ -2098,6 +2116,7 @@ def _db_user_row(*, password: str, password_reset_required: bool | None = None, def _prisma_with_user(row) -> MagicMock: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=row) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=row) return mock_prisma_client @@ -2288,6 +2307,7 @@ class TestScreenLoginPasswordForBreach: assert breached is False mock_prisma_client.db.litellm_usertable.update.assert_not_called() + mock_prisma_client.replica_db = mock_prisma_client.db @pytest.mark.asyncio async def test_rechecks_when_last_check_is_older_than_24_hours(self): @@ -2323,6 +2343,7 @@ class TestScreenLoginPasswordForBreach: assert breached is False mock_prisma_client.db.litellm_usertable.update.assert_not_called() + mock_prisma_client.replica_db = mock_prisma_client.db @pytest.mark.asyncio async def test_db_failure_never_raises_but_still_reports_the_breach(self): @@ -2331,6 +2352,7 @@ class TestScreenLoginPasswordForBreach: password = "Password123!" mock_prisma_client = _prisma_with_user(None) mock_prisma_client.db.litellm_usertable.update = AsyncMock(side_effect=RuntimeError("db down")) + mock_prisma_client.replica_db = mock_prisma_client.db assert ( await screen_login_password_for_breach( diff --git a/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py b/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py index 7506bd031d9..62da254faf6 100644 --- a/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py +++ b/tests/test_litellm/proxy/auth/test_model_access_group_budgets.py @@ -321,6 +321,7 @@ class _RecordingPrismaClient: self.rows = {row.access_group_name: row for row in rows} self.batches: list[list[str]] = [] self.db = SimpleNamespace(litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=self._find_many)) + self.replica_db = self.db async def _find_many(self, **kwargs): requested = list(kwargs["where"]["access_group_name"]["in"]) @@ -498,6 +499,7 @@ async def test_a_database_error_does_not_block_the_request(): class _FailingPrismaClient: def __init__(self) -> None: self.db = SimpleNamespace(litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=self._boom)) + self.replica_db = self.db async def _boom(self, **kwargs): raise RuntimeError("database unavailable") diff --git a/tests/test_litellm/proxy/auth/test_onboarding.py b/tests/test_litellm/proxy/auth/test_onboarding.py index 5d173e57cdf..b6aed050be3 100644 --- a/tests/test_litellm/proxy/auth/test_onboarding.py +++ b/tests/test_litellm/proxy/auth/test_onboarding.py @@ -31,6 +31,7 @@ _POLICY_NO_BREACH_CHECK = {"password_policy_check_breached_passwords": False} class _AsyncTx: def __init__(self, db: MagicMock): self.db = db + self.replica_db = self.db async def __aenter__(self) -> MagicMock: return self.db @@ -63,6 +64,7 @@ def _make_user() -> MagicMock: def _make_prisma(invite: MagicMock, user: MagicMock | None = None) -> MagicMock: prisma = MagicMock() prisma.db.litellm_invitationlink.find_unique = AsyncMock(return_value=invite) + prisma.replica_db = prisma.db prisma.db.litellm_invitationlink.update = AsyncMock() prisma.db.litellm_invitationlink.update_many = AsyncMock(return_value=1) prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user) @@ -125,6 +127,7 @@ async def test_get_token_rejects_already_used_link(): assert "already been used" in exc_info.value.detail["error"] # The user table must never have been queried prisma.db.litellm_usertable.find_unique.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -215,6 +218,7 @@ async def test_get_token_returns_onboarding_token_without_minting_ui_key(): mock_generate_key.assert_not_called() prisma.db.litellm_invitationlink.update_many.assert_not_called() + prisma.replica_db = prisma.db prisma.db.litellm_invitationlink.update.assert_not_called() @@ -247,6 +251,7 @@ async def test_claim_token_rejects_already_used_link(): assert "already been used" in exc_info.value.detail["error"] # Password must never have been written prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -315,6 +320,7 @@ async def test_claim_token_rejects_missing_onboarding_token(): assert exc_info.value.status_code == 401 assert "Missing onboarding session" in exc_info.value.detail["error"] prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -344,6 +350,7 @@ async def test_claim_token_rejects_wrong_onboarding_session(): assert exc_info.value.status_code == 401 assert "Invalid onboarding session" in exc_info.value.detail["error"] prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -371,6 +378,7 @@ async def test_claim_token_rejects_invalid_bearer_token(): assert exc_info.value.status_code == 401 assert "Invalid onboarding session" in exc_info.value.detail["error"] prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -381,6 +389,7 @@ async def test_claim_token_rejects_concurrent_reuse_before_password_write(): invite = _make_invite(is_accepted=False) prisma = _make_prisma(invite) prisma.db.litellm_invitationlink.update_many = AsyncMock(return_value=0) + prisma.replica_db = prisma.db request = _make_claim_request(_make_onboarding_token()) data = InvitationClaim( invitation_link="invite-abc", @@ -456,6 +465,7 @@ async def test_claim_token_sets_accepted_at_after_password_written(): # Password was written prisma.db.litellm_invitationlink.update_many.assert_called_once() + prisma.replica_db = prisma.db reserve_kwargs = prisma.db.litellm_invitationlink.update_many.call_args.kwargs assert reserve_kwargs["where"] == {"id": "invite-abc", "is_accepted": False} assert reserve_kwargs["data"]["is_accepted"] is True @@ -627,6 +637,7 @@ async def test_claim_token_rejects_short_password_before_consuming_invite(): assert exc_info.value.code == "400" assert "at least 12 characters" in exc_info.value.message prisma.db.litellm_invitationlink.update_many.assert_not_called() + prisma.replica_db = prisma.db prisma.db.litellm_usertable.update.assert_not_called() @@ -661,6 +672,7 @@ async def test_claim_token_rejects_breached_password_before_consuming_invite(): assert exc_info.value.code == "400" assert "data breaches" in exc_info.value.message prisma.db.litellm_invitationlink.update_many.assert_not_called() + prisma.replica_db = prisma.db prisma.db.litellm_usertable.update.assert_not_called() @@ -707,3 +719,4 @@ async def test_claim_token_fails_open_when_hibp_unreachable(): assert "token" in result prisma.db.litellm_usertable.update.assert_called_once() + prisma.replica_db = prisma.db diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 2d597abf3b8..a7abcf4d3b2 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -1182,6 +1182,7 @@ def assert_ownership_registered_for_team_a(prisma_client: AsyncMock, batch_id: s assert created["created_by"] == "user_a" assert created["team_id"] == "team_a" prisma_client.db.litellm_managedobjecttable.update_many.assert_not_awaited() + prisma_client.replica_db = prisma_client.db @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py index 0f64ef2b4ca..6e542e9cdaf 100644 --- a/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_config_sync_pubsub.py @@ -683,6 +683,7 @@ async def test_model_repository_write_publishes_via_live_coordination_cache() -> client = _RecordingRedisClient() prisma_client = MagicMock() prisma_client.db.litellm_proxymodeltable.update = AsyncMock(return_value={"model_id": "m-1"}) + prisma_client.replica_db = prisma_client.db repository = ModelRepository(prisma_client) table = repository.table assert isinstance(table, _PublishOnWriteActions) @@ -711,6 +712,7 @@ async def test_ui_settings_write_publishes_via_live_coordination_cache() -> None client = _RecordingRedisClient() prisma_client = MagicMock() prisma_client.db.litellm_uisettings.upsert = AsyncMock(return_value={"id": "ui_settings"}) + prisma_client.replica_db = prisma_client.db table = UISettingsRepository(prisma_client).table assert isinstance(table, _PublishOnWriteActions) @@ -790,6 +792,7 @@ def _reload_config_prisma_client() -> MagicMock: prisma_client = MagicMock() prisma_client.get_generic_data = AsyncMock(return_value=config_record) prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=config_record) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_config.upsert = AsyncMock(return_value=config_record) prisma_client.db.litellm_config.update_many = AsyncMock(return_value=1) return prisma_client @@ -823,6 +826,7 @@ async def test_model_cost_map_reload_does_not_publish_config_change() -> None: _set_redis_usage_cache(previous_cache) prisma_client.db.litellm_config.update_many.assert_awaited_once() + prisma_client.replica_db = prisma_client.db assert client.published == [] @@ -844,6 +848,7 @@ async def test_anthropic_beta_headers_reload_does_not_publish_config_change() -> _set_redis_usage_cache(previous_cache) prisma_client.db.litellm_config.upsert.assert_awaited_once() + prisma_client.replica_db = prisma_client.db assert client.published == [] diff --git a/tests/test_litellm/proxy/common_utils/test_expired_ui_session_key_cleanup_manager.py b/tests/test_litellm/proxy/common_utils/test_expired_ui_session_key_cleanup_manager.py index 8623d93c0a3..3a0f0e4442e 100644 --- a/tests/test_litellm/proxy/common_utils/test_expired_ui_session_key_cleanup_manager.py +++ b/tests/test_litellm/proxy/common_utils/test_expired_ui_session_key_cleanup_manager.py @@ -44,6 +44,7 @@ class TestExpiredUISessionKeyCleanupManager: mock_prisma_client.db.litellm_verificationtoken.find_many.return_value = ( mock_keys ) + mock_prisma_client.replica_db = mock_prisma_client.db with patch( "litellm.proxy.common_utils.expired_ui_session_key_cleanup_manager.datetime" diff --git a/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py b/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py index dd6c1637cad..7ddf3c4dd52 100644 --- a/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py +++ b/tests/test_litellm/proxy/common_utils/test_key_rotation_e2e.py @@ -170,6 +170,7 @@ class TestKeyRotationErrorResilience: be attempted. No key should be silently skipped. """ mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db manager = KeyRotationManager(mock_prisma) key1 = LiteLLM_VerificationToken( @@ -221,6 +222,7 @@ class TestKeyRotationErrorResilience: but process_rotations should catch it per-key. """ mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db manager = KeyRotationManager(mock_prisma) key = LiteLLM_VerificationToken( @@ -260,6 +262,7 @@ class TestKeyRotationErrorResilience: update for rotation_count should still have succeeded (it runs before the hook). """ mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db manager = KeyRotationManager(mock_prisma) key = LiteLLM_VerificationToken( @@ -359,6 +362,7 @@ class TestKeyRotationFullFlow: # Mock cleanup mock_prisma.db.litellm_deprecatedverificationtoken.delete_many.return_value = 1 + mock_prisma.replica_db = mock_prisma.db # Mock find keys mock_prisma.db.litellm_verificationtoken.find_many.return_value = [key] @@ -399,6 +403,7 @@ class TestKeyRotationFullFlow: correctly each time: 0 -> 1 -> 2 -> 3 """ mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db manager = KeyRotationManager(mock_prisma) rotation_counts_seen = [] @@ -445,6 +450,7 @@ class TestKeyRotationFullFlow: """ mock_prisma = AsyncMock() mock_prisma.db.litellm_deprecatedverificationtoken.delete_many.return_value = 0 + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_verificationtoken.find_many.return_value = [] mock_lock = MagicMock() @@ -468,6 +474,7 @@ class TestKeyRotationFullFlow: the DB update for rotation metadata should be skipped. """ mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db manager = KeyRotationManager(mock_prisma) key = LiteLLM_VerificationToken( @@ -525,6 +532,7 @@ class TestKeyRotationInitialization: When no pod_lock_manager is provided, it defaults to None. """ mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db manager = KeyRotationManager(mock_prisma) assert manager.pod_lock_manager is None diff --git a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py index 6103a40d6c7..a0750a8b7ee 100644 --- a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py +++ b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py @@ -57,6 +57,7 @@ class TestKeyRotationManagerPassesKeyAlias: mock_prisma.db.litellm_verificationtoken.update = AsyncMock( return_value=mock_key ) + mock_prisma.replica_db = mock_prisma.db # Create mock response mock_response = GenerateKeyResponse( @@ -114,6 +115,7 @@ class TestKeyRotationManagerPassesKeyAlias: mock_prisma.db.litellm_verificationtoken.update = AsyncMock( return_value=mock_key ) + mock_prisma.replica_db = mock_prisma.db mock_response = GenerateKeyResponse( key="sk-new-key-value", diff --git a/tests/test_litellm/proxy/common_utils/test_key_rotation_manager.py b/tests/test_litellm/proxy/common_utils/test_key_rotation_manager.py index 40a186a9059..a8b18b24490 100644 --- a/tests/test_litellm/proxy/common_utils/test_key_rotation_manager.py +++ b/tests/test_litellm/proxy/common_utils/test_key_rotation_manager.py @@ -30,6 +30,7 @@ class TestKeyRotationManager: """ # Setup mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db manager = KeyRotationManager(mock_prisma_client) now = datetime.now(timezone.utc) @@ -101,6 +102,7 @@ class TestKeyRotationManager: """ # Setup mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db manager = KeyRotationManager(mock_prisma_client) # Use a fixed timestamp to avoid timing issues in tests @@ -171,6 +173,7 @@ class TestKeyRotationManager: """ # Setup mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db manager = KeyRotationManager(mock_prisma_client) # Mock key to rotate @@ -231,6 +234,7 @@ class TestKeyRotationManager: Test that _cleanup_expired_deprecated_keys deletes expired deprecated keys. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_deprecatedverificationtoken.delete_many.return_value = ( 3 ) @@ -251,6 +255,7 @@ class TestKeyRotationManager: Test that _rotate_key passes grace_period in RegenerateKeyRequest. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db manager = KeyRotationManager(mock_prisma_client) key_to_rotate = LiteLLM_VerificationToken( diff --git a/tests/test_litellm/proxy/common_utils/test_periodic_reload_schedule.py b/tests/test_litellm/proxy/common_utils/test_periodic_reload_schedule.py index cabb7452d1a..c5569c13a10 100644 --- a/tests/test_litellm/proxy/common_utils/test_periodic_reload_schedule.py +++ b/tests/test_litellm/proxy/common_utils/test_periodic_reload_schedule.py @@ -33,6 +33,7 @@ def _row(param_value=None, reload_revision=0, last_run_at=None): def _mock_prisma(row=None, upserted_revision=1): prisma_client = MagicMock() prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=row) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_config.upsert = AsyncMock(return_value=_row(reload_revision=upserted_revision)) prisma_client.db.litellm_config.update_many = AsyncMock(return_value=1) return prisma_client @@ -85,6 +86,7 @@ class _FakeConfigTable: def _fake_prisma(table): prisma_client = MagicMock() prisma_client.db.litellm_config = table + prisma_client.replica_db = prisma_client.db return prisma_client @@ -291,6 +293,7 @@ async def test_record_reload_run_updates_last_run_without_creating_or_bumping(): kwargs = prisma_client.db.litellm_config.update_many.await_args.kwargs assert kwargs == {"data": {"last_run_at": LAST_RUN}, "where": {"param_name": "model_cost_map_reload_config"}} prisma_client.db.litellm_config.upsert.assert_not_called() + prisma_client.replica_db = prisma_client.db @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py index ca2ff8bcce1..8524fd45fbc 100644 --- a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py +++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py @@ -167,6 +167,7 @@ async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_rep prisma_client.db.litellm_agentstable.find_unique = AsyncMock( return_value=FakeAgentRow(agent_id, "read-through-db-agent") ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) @@ -193,6 +194,7 @@ async def test_get_agent_with_read_through_recovers_agent_by_name(clean_agent_re prisma_client.db.litellm_agentstable.find_unique = AsyncMock( side_effect=[None, FakeAgentRow("read-through-name-lookup-id", agent_name)] ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) @@ -217,6 +219,7 @@ async def test_get_agent_with_read_through_returns_none_for_unknown_agent( prisma_client: Final = MagicMock() prisma_client.db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) @@ -236,6 +239,7 @@ async def test_resync_agents_already_registered_skips_db(clean_agent_registry, m prisma_client.db.litellm_agentstable.find_unique = AsyncMock( return_value=FakeAgentRow(agent_id, "read-through-dedup-agent") ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) @@ -283,6 +287,7 @@ async def test_get_guardrail_with_read_through_recovers_guardrail_created_on_sib prisma_client.db.litellm_guardrailstable.find_first = AsyncMock( return_value=FakeGuardrailRow(guardrail_id, guardrail_name) ) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( side_effect=AssertionError("full-table guardrail scan on read-through miss") ) @@ -311,6 +316,7 @@ async def test_get_guardrail_with_read_through_returns_none_for_unknown_guardrai prisma_client: Final = MagicMock() prisma_client.db.litellm_guardrailstable.find_first = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) @@ -327,6 +333,7 @@ async def test_resync_guardrails_never_loads_non_active_rows(monkeypatch): pending_name: Final = "pending-review-guardrail" prisma_client: Final = MagicMock() prisma_client.db.litellm_guardrailstable.find_first = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) @@ -353,6 +360,7 @@ async def test_resync_guardrails_syncs_under_guardrail_reconcile_lock(monkeypatc prisma_client.db.litellm_guardrailstable.find_first = AsyncMock( return_value=FakeGuardrailRow("lock-scope-guardrail-id", guardrail_name) ) + prisma_client.replica_db = prisma_client.db lock_states: list[bool] = [] def record_sync(guardrail) -> None: @@ -377,6 +385,7 @@ async def test_resync_model_deployments_mutates_router_under_model_reconcile_loc prisma_client: Final = MagicMock() prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[MagicMock()]) + prisma_client.replica_db = prisma_client.db router: Final = MagicMock() router.get_model_list.return_value = [] lock_states: list[bool] = [] @@ -410,6 +419,7 @@ async def test_resync_model_deployments_loads_db_credentials_before_reconciling_ rows: Final = [MagicMock()] prisma_client: Final = MagicMock() prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=rows) + prisma_client.replica_db = prisma_client.db router: Final = MagicMock() router.get_model_list.return_value = [] installed: Final = MagicMock() @@ -451,6 +461,7 @@ async def test_resync_model_deployments_respects_supported_db_objects(monkeypatc prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock( side_effect=AssertionError("db hit for an object type this replica does not load") ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": ["guardrails"]}) @@ -469,6 +480,7 @@ async def test_resync_guardrails_respects_supported_db_objects(monkeypatch): prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( side_effect=AssertionError("db hit for an object type this replica does not load") ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": ["models"]}) @@ -487,6 +499,7 @@ async def test_resync_agents_respects_supported_db_objects(clean_agent_registry, prisma_client.db.litellm_agentstable.find_unique = AsyncMock( side_effect=AssertionError("db hit for an object type this replica does not load") ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) monkeypatch.setattr(proxy_server, "general_settings", {"supported_db_objects": ["models"]}) @@ -508,6 +521,7 @@ async def test_resync_agents_waits_for_agent_reload_and_skips_duplicate_registra prisma_client.db.litellm_agentstable.find_unique = AsyncMock( side_effect=AssertionError("db hit while the agent reload held the reconcile lock") ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) diff --git a/tests/test_litellm/proxy/db/mcp_server/test_db.py b/tests/test_litellm/proxy/db/mcp_server/test_db.py index e2440e49f19..079f066efd0 100644 --- a/tests/test_litellm/proxy/db/mcp_server/test_db.py +++ b/tests/test_litellm/proxy/db/mcp_server/test_db.py @@ -14,6 +14,7 @@ from litellm.proxy._experimental.mcp_server.db import ( def _prisma_client_returning(team_record: object) -> MagicMock: prisma_client = MagicMock() prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_record) + prisma_client.replica_db = prisma_client.db return prisma_client @@ -42,11 +43,13 @@ async def test_fetch_mcp_servers_by_team(team_record, expected): where={"team_id": "team-123"}, include={"object_permission": True}, ) + prisma_client.replica_db = prisma_client.db def _prisma_client_with_missing_mcp_server_row() -> MagicMock: prisma_client = MagicMock() prisma_client.db.litellm_mcpservertable.update = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db return prisma_client diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index ff0b67d426b..b524902329f 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -76,6 +76,7 @@ class _FakePrismaClient: litellm_verificationtoken=_InFlightCountingTable(), litellm_projecttable=_FakeFindUniqueTable(row=project_row), ) + self.replica_db = self.db def _row(window_start: datetime, spend: float) -> SimpleNamespace: @@ -97,7 +98,8 @@ class _PausedSpendTable: async def _reseed_with_paused_table( table: _PausedSpendTable, cache: DualCache, counter_key: str, window: bool ) -> float | None: - prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_usertable=table, litellm_budgetwindowspend=table)) + tables: Final = SimpleNamespace(litellm_usertable=table, litellm_budgetwindowspend=table) + prisma: Final = SimpleNamespace(db=tables, replica_db=tables) if window: return await SpendCounterReseed.coalesced_window( prisma_client=prisma, diff --git a/tests/test_litellm/proxy/db/test_spend_log_tool_index.py b/tests/test_litellm/proxy/db/test_spend_log_tool_index.py index 282c2a7cfaa..4b4af3304c5 100644 --- a/tests/test_litellm/proxy/db/test_spend_log_tool_index.py +++ b/tests/test_litellm/proxy/db/test_spend_log_tool_index.py @@ -40,6 +40,7 @@ class _FakeBatcher: def _prisma(batch_: MagicMock) -> MagicMock: prisma = MagicMock() prisma.db.batch_ = batch_ + prisma.replica_db = prisma.db prisma.db.litellm_spendlogtoolindex.create_many = AsyncMock() return prisma @@ -296,12 +297,14 @@ class TestFlushToolUsageTransactions: ] batcher.litellm_spendlogtoolindex.create_many.assert_not_called() prisma.db.batch_.assert_called_once() + prisma.replica_db = prisma.db assert batcher.litellm_dailytoolspend.upsert.call_count == len(tool_names) @pytest.mark.asyncio async def test_index_connection_error_is_retried_before_the_rollup_is_attempted(self, monkeypatch): prisma, batcher = _prisma_with_batcher() prisma.db.litellm_spendlogtoolindex.create_many = AsyncMock(side_effect=[httpx.ConnectError("down"), None]) + prisma.replica_db = prisma.db async def fake_sleep(seconds: float) -> None: return None @@ -310,12 +313,14 @@ class TestFlushToolUsageTransactions: await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) assert prisma.db.litellm_spendlogtoolindex.create_many.await_count == 2 prisma.db.batch_.assert_called_once() + prisma.replica_db = prisma.db batcher.litellm_dailytoolspend.upsert.assert_called_once() @pytest.mark.asyncio async def test_ambiguous_index_error_drops_the_batch_without_touching_the_rollup(self): prisma, _ = _prisma_with_batcher() prisma.db.litellm_spendlogtoolindex.create_many = AsyncMock(side_effect=httpx.ReadTimeout("ambiguous")) + prisma.replica_db = prisma.db with pytest.raises(httpx.ReadTimeout): await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) prisma.db.litellm_spendlogtoolindex.create_many.assert_awaited_once() @@ -326,6 +331,7 @@ class TestFlushToolUsageTransactions: prisma, _ = _prisma_with_batcher() await flush_tool_usage_transactions(prisma_client=prisma, transactions=[]) prisma.db.batch_.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio async def test_connection_errors_retry_and_succeed(self, monkeypatch): @@ -364,6 +370,7 @@ class TestFlushToolUsageTransactions: with pytest.raises(ValueError, match="bad data"): await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) prisma.db.batch_.assert_called_once() + prisma.replica_db = prisma.db @pytest.mark.asyncio @pytest.mark.parametrize("ambiguous_error", ["ReadTimeout", "ReadError"]) @@ -377,3 +384,4 @@ class TestFlushToolUsageTransactions: with pytest.raises((httpx.ReadTimeout, httpx.ReadError)): await flush_tool_usage_transactions(prisma_client=prisma, transactions=[_transaction("r1")]) prisma.db.batch_.assert_called_once() + prisma.replica_db = prisma.db diff --git a/tests/test_litellm/proxy/db/test_tool_registry_writer.py b/tests/test_litellm/proxy/db/test_tool_registry_writer.py index 6318e4422cf..ba1c2639862 100644 --- a/tests/test_litellm/proxy/db/test_tool_registry_writer.py +++ b/tests/test_litellm/proxy/db/test_tool_registry_writer.py @@ -58,6 +58,7 @@ def _make_prisma( """Return a mock prisma_client with litellm_tooltable.upsert, find_many, find_unique.""" prisma = MagicMock() prisma.db.litellm_tooltable = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_tooltable.upsert = AsyncMock(return_value=upsert_return) prisma.db.litellm_tooltable.find_many = AsyncMock( return_value=find_many_rows if find_many_rows is not None else [] @@ -72,6 +73,7 @@ async def test_batch_upsert_tools_calls_upsert(): items = [{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}] await batch_upsert_tools(prisma, items) prisma.db.litellm_tooltable.upsert.assert_awaited_once() + prisma.replica_db = prisma.db call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs assert call_kw["where"] == {"tool_name": "tool_a"} assert call_kw["data"]["create"]["tool_name"] == "tool_a" @@ -88,6 +90,7 @@ async def test_batch_upsert_tools_empty_list(): prisma = _make_prisma() await batch_upsert_tools(prisma, []) prisma.db.litellm_tooltable.upsert.assert_not_awaited() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -96,6 +99,7 @@ async def test_batch_upsert_tools_skips_empty_names(): items = [{"tool_name": "", "origin": None}, {"tool_name": None}] # type: ignore[list-item] await batch_upsert_tools(prisma, items) prisma.db.litellm_tooltable.upsert.assert_not_awaited() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -128,6 +132,7 @@ async def test_list_tools_no_filter(): assert result[0].tool_name == "tool_a" assert result[0].call_count == 5 prisma.db.litellm_tooltable.find_many.assert_awaited_once() + prisma.replica_db = prisma.db call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs assert call_kw["where"] == {} assert call_kw["order"] == {"created_at": "desc"} @@ -161,6 +166,7 @@ async def test_get_tool_found(): prisma.db.litellm_tooltable.find_unique.assert_awaited_once_with( where={"tool_name": "my_tool"} ) + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -185,6 +191,7 @@ async def test_update_tool_policy_calls_upsert_then_get_tool(): assert result is not None assert result.input_policy == "blocked" prisma.db.litellm_tooltable.upsert.assert_awaited_once() + prisma.replica_db = prisma.db call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs assert call_kw["where"] == {"tool_name": "my_tool"} assert call_kw["data"]["update"]["input_policy"] == "blocked" @@ -213,6 +220,7 @@ async def test_get_tools_by_names_returns_policy_map(): prisma.db.litellm_tooltable.find_many.assert_awaited_once_with( where={"tool_name": {"in": ["tool_a", "tool_b"]}} ) + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -221,6 +229,7 @@ async def test_get_tools_by_names_empty_list(): result = await get_tools_by_names(prisma, []) assert result == {} prisma.db.litellm_tooltable.find_many.assert_not_awaited() + prisma.replica_db = prisma.db # --- ToolPolicyRegistry --- @@ -256,6 +265,7 @@ async def test_tool_policy_registry_sync_and_get_effective_policies(): _mock_tool_row("tool_c", input_policy="untrusted"), ] ) + prisma.replica_db = prisma.db prisma.db.litellm_objectpermissiontable.find_many = AsyncMock( return_value=[ _mock_perm_row("op-key-1", ["tool_a"]), @@ -310,6 +320,7 @@ async def test_sync_tool_policy_from_db_retries_on_transport_error_first_read(): mock_prisma_client.db.litellm_tooltable.find_many = AsyncMock( side_effect=_flaky_find_many ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable.find_many = AsyncMock( return_value=[] ) @@ -346,6 +357,7 @@ async def test_sync_tool_policy_from_db_retries_on_transport_error_second_read() mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_tooltable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable.find_many = AsyncMock( side_effect=_flaky_perms_find_many ) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index bf641fd6cd0..09091f4bd80 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -86,6 +86,7 @@ def mock_prisma_client(mocker): mock_client = mocker.Mock() # Create async mocks for the database methods mock_client.db = mocker.Mock() + mock_client.replica_db = mock_client.db mock_client.db.litellm_guardrailstable = mocker.Mock() mock_client.db.litellm_guardrailstable.find_many = AsyncMock( return_value=[MOCK_DB_GUARDRAIL] @@ -175,6 +176,7 @@ async def test_list_guardrails_v2_skips_stale_db_backed_in_memory_entries(mocker } mock_prisma_client = mocker.Mock() mock_prisma_client.db = mocker.Mock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[]) @@ -211,6 +213,7 @@ async def test_get_guardrail_info_404s_stale_db_backed_entry( mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( return_value=None ) + mock_prisma_client.replica_db = mock_prisma_client.db # In-memory still has it, but it's tagged as 'db' (stale, awaiting reconcile) mock_in_memory_handler.get_source.return_value = "db" @@ -240,6 +243,7 @@ async def test_list_guardrails_v2_masks_sensitive_data_in_db_guardrails(mocker): mock_prisma_client = mocker.Mock() mock_prisma_client.db = mocker.Mock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( return_value=[db_guardrail_with_secrets] @@ -295,6 +299,7 @@ async def test_list_guardrails_v2_masks_sensitive_data_in_config_guardrails(mock mock_prisma_client = mocker.Mock() mock_prisma_client.db = mocker.Mock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[]) @@ -354,6 +359,7 @@ async def test_list_guardrails_v2_admin_viewer_sees_guardrails_of_teams_they_are mock_prisma_client = mocker.Mock() mock_prisma_client.db = mocker.Mock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( return_value=[other_team_guardrail] @@ -403,6 +409,7 @@ async def test_list_guardrails_v2_masks_sensitive_data_for_admin_viewer(mocker): mock_prisma_client = mocker.Mock() mock_prisma_client.db = mocker.Mock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_guardrailstable = mocker.Mock() mock_prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( return_value=[other_team_guardrail_with_secrets] @@ -465,6 +472,7 @@ async def test_get_guardrail_info_from_config( mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( return_value=None ) + mock_prisma_client.replica_db = mock_prisma_client.db response = await get_guardrail_info("test-config-guardrail") @@ -489,6 +497,7 @@ async def test_get_guardrail_info_not_found( mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( return_value=None ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_in_memory_handler.get_guardrail_by_id.return_value = None with pytest.raises(HTTPException) as exc_info: @@ -1903,6 +1912,7 @@ async def test_register_guardrail_success(mocker): """Register creates a row with status pending_review and returns guardrail_id.""" mock_prisma = mocker.Mock() mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db created_row = mocker.Mock( guardrail_id="reg-123", guardrail_name=MOCK_REGISTER_REQUEST.guardrail_name, @@ -1961,6 +1971,7 @@ async def test_register_guardrail_non_admin_cross_team_allowed(mocker): """Non-admin may register for a team in their user.teams list even if the key's team_id differs.""" mock_prisma = mocker.Mock() mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db created = mocker.Mock( guardrail_id="g1", guardrail_name=MOCK_REGISTER_REQUEST.guardrail_name, @@ -2016,6 +2027,7 @@ async def test_register_guardrail_duplicate_name(mocker): mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock( return_value={"guardrail_name": MOCK_REGISTER_REQUEST.guardrail_name} ) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_id="u1", user_email="a@b.com", team_id="team-1") @@ -2043,6 +2055,7 @@ async def test_list_guardrail_submissions_non_admin_scoped_to_own_teams(mocker): ) find_many = AsyncMock(return_value=[own_team_row]) mock_prisma.db.litellm_guardrailstable.find_many = find_many + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) mocker.patch( "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", @@ -2068,6 +2081,7 @@ async def test_list_guardrail_submissions_non_admin_no_teams(mocker): mock_prisma = mocker.Mock() find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_guardrailstable.find_many = find_many + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) mocker.patch( "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", @@ -2121,6 +2135,7 @@ async def test_list_guardrail_submissions_success(mocker): updated_at=datetime.now(), ) mock_prisma.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[row]) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2140,6 +2155,7 @@ async def test_list_guardrail_submissions_returns_only_team_guardrails(mocker): mock_prisma = mocker.Mock() find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_guardrailstable.find_many = find_many + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2181,6 +2197,7 @@ async def test_list_guardrail_submissions_team_id_filter(mocker): ) find_many = AsyncMock(return_value=[row_abc, row_other]) mock_prisma.db.litellm_guardrailstable.find_many = find_many + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2199,6 +2216,7 @@ async def test_get_guardrail_submission_not_found(mocker): """Get submission returns 404 when guardrail_id does not exist.""" mock_prisma = mocker.Mock() mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2224,6 +2242,7 @@ async def test_get_guardrail_submission_non_admin_own_team(mocker): updated_at=datetime.now(), ) mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) mocker.patch( "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", @@ -2254,6 +2273,7 @@ async def test_get_guardrail_submission_non_admin_other_team_forbidden(mocker): updated_at=datetime.now(), ) mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) mocker.patch( "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", @@ -2283,6 +2303,7 @@ async def test_get_guardrail_submission_admin_viewer_other_team_allowed(mocker): updated_at=datetime.now(), ) mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) mock_get_user_team_ids = mocker.patch( "litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids", @@ -2315,6 +2336,7 @@ async def test_approve_guardrail_submission_success(mocker): guardrail_info={}, ) mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_guardrailstable.update = AsyncMock() mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) mock_handler = mocker.Mock() @@ -2340,6 +2362,7 @@ async def test_approve_guardrail_submission_not_pending(mocker): mock_prisma = mocker.Mock() row = mocker.Mock(guardrail_id="x", guardrail_name="y", status="active") mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2354,6 +2377,7 @@ async def test_reject_guardrail_submission_success(mocker): mock_prisma = mocker.Mock() row = mocker.Mock(guardrail_id="rej-1", guardrail_name="r", status="pending_review") mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_guardrailstable.update = AsyncMock() mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2374,6 +2398,7 @@ async def test_reject_guardrail_submission_not_pending(mocker): guardrail_id="already-active", guardrail_name="g", status="active" ) mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2430,6 +2455,7 @@ async def test_register_guardrail_accepts_valid_https_url(mocker): """Register accepts valid https api_base URLs.""" mock_prisma = mocker.Mock() mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db created_row = mocker.Mock( guardrail_id="valid-url-123", guardrail_name="valid-guard", @@ -2470,6 +2496,7 @@ async def test_approve_guardrail_init_failure_returns_warning(mocker): guardrail_info={}, ) mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_guardrailstable.update = AsyncMock() mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -2507,6 +2534,7 @@ async def test_approve_guardrail_no_warning_on_success(mocker): guardrail_info={}, ) mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_guardrailstable.update = AsyncMock() mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -2530,6 +2558,7 @@ async def test_list_submissions_single_db_query(mocker): mock_prisma = mocker.Mock() find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_guardrailstable.find_many = find_many + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) @@ -2568,6 +2597,7 @@ async def test_list_submissions_summary_counts_unaffected_by_filters(mocker): ) all_rows = [pending_row, active_row] mock_prisma.db.litellm_guardrailstable.find_many = AsyncMock(return_value=all_rows) + mock_prisma.replica_db = mock_prisma.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 836668de0c8..457937d89df 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -934,6 +934,7 @@ class TestScanOnlyToolResultsInitRefusal: async def test_update_guardrail_in_db_raises_when_row_missing(): prisma_client = MagicMock() prisma_client.db.litellm_guardrailstable.update = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db with pytest.raises( Exception, diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py index db87e12ac88..b959fd28045 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -111,6 +111,7 @@ def _prisma( ) -> MagicMock: client = MagicMock() db = client.db + client.replica_db = client.db db.litellm_guardrailstable.find_many = AsyncMock(return_value=find_many or []) db.litellm_guardrailstable.find_unique = AsyncMock(return_value=find_unique) db.litellm_dailyguardrailmetrics.find_many = AsyncMock(return_value=metrics or []) @@ -309,6 +310,7 @@ def _units_table_missing() -> TableNotFoundError: async def test_overview_degrades_units_to_empty_when_units_table_is_missing(): prisma = _prisma(metrics=[_metric("yaml-pii", requests=4, passed=3, blocked=1)]) prisma.db.litellm_dailyguardrailusageunits.find_many = AsyncMock(side_effect=_units_table_missing()) + prisma.replica_db = prisma.db handler = _config_handler(_yaml_guardrail(guardrail_id="yaml-uuid", name="yaml-pii")) p1, p2 = _patches(prisma, handler) with p1, p2: @@ -443,6 +445,7 @@ async def test_detail_breaks_cost_down_by_unit_day_team_and_key(): async def test_detail_degrades_units_to_empty_when_units_table_is_missing(): prisma = _prisma(metrics=[_metric("yaml-pii", requests=4, passed=3, blocked=1)]) prisma.db.litellm_dailyguardrailusageunits.find_many = AsyncMock(side_effect=_units_table_missing()) + prisma.replica_db = prisma.db handler = _config_handler(_yaml_guardrail()) p1, p2 = _patches(prisma, handler) with p1, p2: @@ -526,6 +529,7 @@ async def test_logs_reports_flagged_action_for_guardrail_flagged_status(): _spend_log("r-block", "guardrail_intervened"), ] ) + prisma.replica_db = prisma.db p1, p2 = _patches(prisma, _config_handler()) with p1, p2: resp = await guardrails_usage_logs( @@ -565,6 +569,7 @@ async def test_logs_reports_post_call_flag_when_pre_call_allowed(): prisma.db.litellm_spendlogs.find_many = AsyncMock( return_value=[_spend_log("r-post-flag", "success", "guardrail_flagged")] ) + prisma.replica_db = prisma.db p1, p2 = _patches(prisma, _config_handler()) with p1, p2: resp = await guardrails_usage_logs( @@ -649,6 +654,7 @@ async def test_policies_overview_returns_a_full_row_and_totals(): metric.policy_id = "pol-1" prisma = _prisma() prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[policy]) + prisma.replica_db = prisma.db prisma.db.litellm_dailypolicymetrics.find_many = AsyncMock(return_value=[metric]) p1, p2 = _patches(prisma, _config_handler()) with p1, p2: @@ -702,6 +708,7 @@ async def test_logs_report_not_run_entries_as_not_run_not_passed(): } prisma = _prisma(find_unique=_db_row(), index_find_many=[index_row]) prisma.db.litellm_spendlogs.find_many = AsyncMock(return_value=[spend_log]) + prisma.replica_db = prisma.db handler = _config_handler() p1, p2 = _patches(prisma, handler) with p1, p2: @@ -732,6 +739,7 @@ async def test_logs_action_passed_filter_excludes_not_run_entries(): spend_log.metadata = {"guardrail_information": [{"guardrail_name": "db-1", "guardrail_status": "not_run"}]} prisma = _prisma(find_unique=_db_row(), index_find_many=[index_row]) prisma.db.litellm_spendlogs.find_many = AsyncMock(return_value=[spend_log]) + prisma.replica_db = prisma.db handler = _config_handler() p1, p2 = _patches(prisma, handler) with p1, p2: diff --git a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py index 69ec098b840..ecc8cb5473a 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py @@ -18,6 +18,7 @@ from litellm.proxy.guardrails.usage_tracking import ( def _prisma() -> MagicMock: client = MagicMock() db = client.db + client.replica_db = client.db db.litellm_dailyguardrailmetrics.upsert = AsyncMock() db.litellm_dailyguardrailusageunits.upsert = AsyncMock() db.litellm_spendlogguardrailindex.create_many = AsyncMock() @@ -142,6 +143,7 @@ async def test_one_failing_upsert_does_not_drop_remaining_writes(): """ prisma = _prisma() prisma.db.litellm_dailyguardrailmetrics.upsert.side_effect = httpx.ConnectError("db down") + prisma.replica_db = prisma.db prisma.db.litellm_dailyguardrailusageunits.upsert.side_effect = [httpx.ConnectError("db down"), None, None] sleep, _ = _fake_sleep() logs = [ @@ -167,6 +169,7 @@ async def test_transient_upsert_failure_is_retried_with_backoff_for_failed_rows_ """ prisma = _prisma() prisma.db.litellm_dailyguardrailusageunits.upsert.side_effect = [httpx.ConnectError("blip"), None, None] + prisma.replica_db = prisma.db sleep, delays = _fake_sleep() logs = [ _payload("r1", usage={"topicPolicyUnits": 1}), @@ -185,6 +188,7 @@ async def test_transient_upsert_failure_is_retried_with_backoff_for_failed_rows_ async def test_persistent_upsert_failure_stops_after_three_retries(): prisma = _prisma() prisma.db.litellm_dailyguardrailmetrics.upsert.side_effect = httpx.ConnectError("db down") + prisma.replica_db = prisma.db sleep, delays = _fake_sleep() pending = PendingRollups() @@ -215,6 +219,7 @@ async def test_retry_exhausted_rows_are_requeued_and_land_on_the_next_flush(): pending = PendingRollups() down = _prisma() down.db.litellm_dailyguardrailmetrics.upsert.side_effect = httpx.ConnectError("db down") + down.replica_db = down.db down.db.litellm_dailyguardrailusageunits.upsert.side_effect = httpx.ConnectError("db down") sleep, _ = _fake_sleep() @@ -249,6 +254,7 @@ async def test_ambiguous_failures_are_never_requeued(): pending = PendingRollups() prisma = _prisma() prisma.db.litellm_dailyguardrailusageunits.upsert.side_effect = httpx.ReadTimeout("maybe committed") + prisma.replica_db = prisma.db sleep, delays = _fake_sleep() await process_spend_logs_guardrail_usage( @@ -296,6 +302,7 @@ async def test_post_send_failure_is_never_retried_so_increments_cannot_double_co httpx.ConnectError("refused"), None, ] + prisma.replica_db = prisma.db sleep, delays = _fake_sleep() logs = [ _payload("r1", usage={"topicPolicyUnits": 1}), @@ -314,6 +321,7 @@ async def test_post_send_failure_is_never_retried_so_increments_cannot_double_co async def test_generic_upsert_exception_is_terminal_for_that_row_only(): prisma = _prisma() prisma.db.litellm_dailyguardrailmetrics.upsert.side_effect = RuntimeError("constraint violation") + prisma.replica_db = prisma.db sleep, delays = _fake_sleep() await process_spend_logs_guardrail_usage(prisma, [_payload("r1", usage={"topicPolicyUnits": 1})], sleep=sleep) @@ -565,6 +573,7 @@ async def test_requeued_cost_is_added_to_the_next_flush(): pending = PendingRollups() down = _prisma() down.db.litellm_dailyguardrailmetrics.upsert.side_effect = httpx.ConnectError("db down") + down.replica_db = down.db down.db.litellm_dailyguardrailusageunits.upsert.side_effect = httpx.ConnectError("db down") sleep, _ = _fake_sleep() @@ -647,6 +656,7 @@ async def test_one_failing_index_statement_does_not_drop_the_others_or_the_rollu monkeypatch.setattr(usage_tracking, "SPEND_LOG_WRITE_BATCH_MAX_ROWS", 100) prisma = _prisma() prisma.db.litellm_spendlogguardrailindex.create_many.side_effect = [None, httpx.ReadTimeout("ambiguous"), None] + prisma.replica_db = prisma.db guardrail_ids = tuple(f"guard-{i}" for i in range(50)) logs = [_fan_out_payload(f"r{i}", guardrail_ids) for i in range(5)] diff --git a/tests/test_litellm/proxy/hooks/test_user_management_event_hooks.py b/tests/test_litellm/proxy/hooks/test_user_management_event_hooks.py index 102430ab985..d6ef29a3d19 100644 --- a/tests/test_litellm/proxy/hooks/test_user_management_event_hooks.py +++ b/tests/test_litellm/proxy/hooks/test_user_management_event_hooks.py @@ -26,6 +26,7 @@ class FakeUserTable: class FakePrismaClient: def __init__(self, rows: List[Dict[str, Any]]): self.db = SimpleNamespace(litellm_usertable=FakeUserTable(rows)) + self.replica_db = self.db async def _run_created_hook(prisma_client: FakePrismaClient, audit_log: AsyncMock) -> None: diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py index d35a676a28b..80adce26e74 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py @@ -38,6 +38,7 @@ def _build_prisma_with_keys(user_keys, mock_user=None, updated_user=None): mock_client = MagicMock() mock_db = MagicMock() mock_client.db = mock_db + mock_client.replica_db = mock_client.db if mock_user is not None: mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) if updated_user is not None: diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py index edcce16ab41..9031bf8ca6b 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py @@ -40,6 +40,7 @@ async def test_patch_user_updates_fields(): mock_client = MagicMock() mock_db = MagicMock() mock_client.db = mock_db + mock_client.replica_db = mock_client.db mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update) mock_db.litellm_teamtable.find_unique = AsyncMock(return_value=None) @@ -105,6 +106,7 @@ async def test_patch_user_manages_group_memberships(): mock_client = MagicMock() mock_db = MagicMock() mock_client.db = mock_db + mock_client.replica_db = mock_client.db mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update) mock_db.litellm_teamtable.find_unique = AsyncMock(return_value=None) @@ -196,6 +198,7 @@ async def test_patch_user_deprovision_without_path(): mock_client = MagicMock() mock_db = MagicMock() mock_client.db = mock_db + mock_client.replica_db = mock_client.db mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update) mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -275,6 +278,7 @@ async def test_patch_user_multiple_fields_without_path(): mock_client = MagicMock() mock_db = MagicMock() mock_client.db = mock_db + mock_client.replica_db = mock_client.db mock_db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_db.litellm_usertable.update = AsyncMock(side_effect=mock_update) mock_db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py index 135175dd29d..5ff3713333b 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py @@ -85,6 +85,7 @@ def mock_prisma_client(): mock_client = MagicMock() mock_db = MagicMock() mock_client.db = mock_db + mock_client.replica_db = mock_client.db mock_find_unique = AsyncMock() mock_db.litellm_teamtable.find_unique = mock_find_unique diff --git a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py index 70e9a96b316..846076f94bf 100644 --- a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py +++ b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py @@ -629,6 +629,7 @@ async def test_get_all_search_tools_from_db_retries_on_transport_error(): mock_prisma_client.db.litellm_searchtoolstable.find_many = AsyncMock( side_effect=_flaky_find_many ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py index 13e57408bd0..d63841630aa 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py @@ -149,6 +149,7 @@ def client_and_mocks(monkeypatch): tx=mock_tx, ) mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) @@ -208,6 +209,7 @@ def test_create_access_group_success(client_and_mocks, base_path, payload): """Create access group with various payloads returns 201.""" client, mock_prisma, mock_table, *_ = client_and_mocks mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[_make_team_record("team-1")]) + mock_prisma.replica_db = mock_prisma.db resp = client.post(base_path, json=payload) assert resp.status_code == 201 @@ -303,6 +305,7 @@ def test_list_access_groups_success_empty(client_and_mocks, base_path): assert resp.json() == [] mock_table.find_many.assert_awaited_once() mock_prisma.db.litellm_teamtable.find_many.assert_not_awaited() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS) @@ -454,6 +457,7 @@ def test_get_access_group_empty_column_and_no_teams_returns_empty(client_and_moc mock_table.find_unique = AsyncMock(return_value=_make_access_group_record(access_group_id="ag-123")) mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db resp = client.get("/v1/access_group/ag-123") assert resp.status_code == 200 @@ -1436,6 +1440,7 @@ def test_update_access_group_null_assigned_ids_treated_as_empty(client_and_mocks def _mock_resource_tables(mock_prisma, *, mcp_servers=(), agents=(), teams=(), keys=()): mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(mcp_servers)) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=list(agents)) mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=list(teams)) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=list(keys)) @@ -1550,6 +1555,7 @@ def test_list_access_groups_skips_lookups_when_nothing_to_resolve(client_and_moc assert all(group["access_mcp_servers"] == [] and group["assigned_keys"] == [] for group in resp.json()) mock_prisma.db.litellm_mcpservertable.find_many.assert_not_awaited() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_many.assert_not_awaited() mock_prisma.db.litellm_verificationtoken.find_many.assert_not_awaited() @@ -1559,6 +1565,7 @@ def test_create_access_group_response_carries_resolved_names(client_and_mocks): client, mock_prisma, *_ = client_and_mocks team_record = _make_team_record("team-1", team_alias="Platform") mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_record) + mock_prisma.replica_db = mock_prisma.db _mock_resource_tables( mock_prisma, mcp_servers=[_make_mcp_server_record("mcp-a", alias="GitHub")], diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index 59c2921e0d0..5fc9905b02e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -58,6 +58,7 @@ async def test_create_duplicate_access_group_fails(): ) ] ) + mock_prisma.replica_db = mock_prisma.db mock_user = UserAPIKeyAuth( user_id="test_admin", @@ -103,6 +104,7 @@ async def test_create_access_group_with_model_ids_tags_only_specific_deployments mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( return_value=deploy_a ) @@ -172,6 +174,7 @@ async def test_create_access_group_with_model_names_tags_all_deployments(): mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock( side_effect=[[], [deploy_a, deploy_b, deploy_c]] ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() mock_user = UserAPIKeyAuth( @@ -217,6 +220,7 @@ async def test_create_access_group_model_ids_takes_priority_over_model_names(): mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( return_value=deploy_a ) @@ -298,6 +302,7 @@ async def test_create_access_group_invalid_model_id_returns_400(): mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) mock_user = UserAPIKeyAuth( @@ -342,6 +347,7 @@ async def test_create_access_group_surfaces_dropped_models(): mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=deploy_a) mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() @@ -385,6 +391,7 @@ async def test_create_access_group_trusts_reload_snapshot_over_post_lock_fresh_r mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=deploy_a) mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() @@ -421,6 +428,7 @@ async def test_tag_deployment_parses_string_model_info_and_refuses_corrupt(): mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + mock_prisma.replica_db = mock_prisma.db pair = await _tag_deployment_with_access_group( model_id="deploy-str", @@ -452,6 +460,7 @@ async def test_delete_access_group_ignores_models_that_were_already_dead(): mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[deploy_broken]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() mock_prisma.db.litellm_modelaccessgroupbudgettable.delete = AsyncMock(return_value=None) @@ -514,6 +523,7 @@ async def test_create_access_group_read_through_recovers_model_created_on_siblin mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=[[db_row], [], [db_row]]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() with ( @@ -560,6 +570,7 @@ async def test_create_access_group_model_missing_everywhere_still_400s(): ) mock_prisma = MagicMock() mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db with ( patch("litellm.proxy.proxy_server.llm_router", mock_router), @@ -714,6 +725,7 @@ class _FakePrismaClient: litellm_modelaccessgroupbudgettable=self.access_group_budget_table, litellm_proxymodeltable=self.model_table, ) + self.replica_db = self.db def jsonify_object(self, data): return dict(data) diff --git a/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py b/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py index 61583d11dfa..5de0fb3ca36 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py +++ b/tests/test_litellm/proxy/management_endpoints/test_activity_tenant_scoping.py @@ -66,6 +66,7 @@ async def test_team_activity_requires_admin_on_every_requested_team(): _make_team("team-B", admin_user_ids=["bob"]), ] ) + prisma.replica_db = prisma.db user_keys = MagicMock(token="alice-key-1") prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[user_keys]) @@ -123,6 +124,7 @@ async def test_team_activity_full_view_when_admin_of_all_requested_teams(): _make_team("team-B", admin_user_ids=["alice"]), ] ) + prisma.replica_db = prisma.db user_info = MagicMock() user_info.teams = ["team-A", "team-B"] @@ -171,6 +173,7 @@ async def test_agent_activity_admin_unscoped(): prisma = MagicMock() prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) + prisma.replica_db = prisma.db captured = {} @@ -216,6 +219,7 @@ async def test_agent_activity_non_admin_no_perms_falls_back_to_owned(): # First call: lookup of owned agents (created_by=alice). # Second call: agent_metadata fetch for the resolved set. prisma.db.litellm_agentstable.find_many = AsyncMock(side_effect=[owned, owned]) + prisma.replica_db = prisma.db captured = {} @@ -264,6 +268,7 @@ async def test_agent_activity_non_admin_intersects_explicit_agent_ids(): prisma = MagicMock() prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) + prisma.replica_db = prisma.db captured = {} @@ -314,6 +319,7 @@ async def test_agent_activity_keyless_caller_does_not_query_created_by_null(): prisma = MagicMock() prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) + prisma.replica_db = prisma.db fake_get_daily = AsyncMock() @@ -359,6 +365,7 @@ async def test_agent_activity_non_admin_no_access_returns_empty_page(): prisma = MagicMock() prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) + prisma.replica_db = prisma.db fake_get_daily = AsyncMock() diff --git a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py index 2f3be61d00f..409cd6bf599 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py @@ -26,6 +26,7 @@ def client_and_mocks(monkeypatch): litellm_budgettable=mock_table, litellm_dailyspend=mock_table, ) + mock_prisma.replica_db = mock_prisma.db # Monkeypatch Mocked Prisma client into the server module monkeypatch.setattr(ps, "prisma_client", mock_prisma) diff --git a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py index 9a2dd914866..90c0ae092e8 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py @@ -198,6 +198,7 @@ async def test_update_cache_settings_persists_url_precedence(monkeypatch): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() @@ -260,6 +261,7 @@ async def test_get_cache_settings_masks_password_bearing_url(): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) + mock_prisma.replica_db = mock_prisma.db proxy_config = MagicMock() proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) @@ -385,6 +387,7 @@ class TestCacheSettingsManager: mock_cache_config = MagicMock() mock_cache_config.cache_settings = '{"type": "redis", "host": "localhost", "port": "6379"}' mock_prisma_client.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=mock_cache_config) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock proxy_config mock_proxy_config = MagicMock() @@ -430,6 +433,7 @@ class TestCacheSettingsManager: mock_cache_config = MagicMock() mock_cache_config.cache_settings = '{"type": "redis", "host": "localhost", "port": "6379"}' mock_prisma_client.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=mock_cache_config) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock proxy_config mock_proxy_config = MagicMock() @@ -468,6 +472,7 @@ class TestCacheSettingsManager: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_cacheconfig.find_unique = AsyncMock(side_effect=_flaky_find_unique) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 @@ -502,6 +507,7 @@ async def test_update_cache_settings_emits_audit_log_when_enabled(monkeypatch): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() @@ -572,6 +578,7 @@ async def test_update_cache_settings_no_audit_when_disabled(monkeypatch): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = MagicMock() @@ -800,6 +807,7 @@ async def test_get_cache_settings_falls_back_to_redis_env(monkeypatch): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db proxy_config = MagicMock() proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) @@ -828,6 +836,7 @@ async def test_get_cache_settings_redacts_password_with_marker(monkeypatch): ) mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) + mock_prisma.replica_db = mock_prisma.db proxy_config = MagicMock() proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) @@ -857,6 +866,7 @@ async def test_get_cache_settings_url_mode_hides_env_discrete_fields(monkeypatch cache_row.cache_settings = {"type": "redis", "url": "redis://:pw@stored-host:6379/0"} mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) + mock_prisma.replica_db = mock_prisma.db proxy_config = MagicMock() proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) @@ -896,6 +906,7 @@ async def test_update_preserves_stored_password_on_redacted_resubmit(monkeypatch existing.cache_settings = {"type": "redis", "host": "oldhost", "password": "realpw"} mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = _mock_proxy_config_identity_crypto() @@ -929,6 +940,7 @@ async def test_update_drops_env_sourced_redacted_secret(monkeypatch): mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = _mock_proxy_config_identity_crypto() @@ -958,6 +970,7 @@ async def test_update_applies_new_password(monkeypatch): existing.cache_settings = json.dumps({"type": "redis", "host": "h", "password": "oldpw"}) mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() proxy_config = _mock_proxy_config_identity_crypto() @@ -1028,6 +1041,7 @@ async def test_get_cache_settings_does_not_surface_non_display_env_credentials(m mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db proxy_config = MagicMock() proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) @@ -1057,6 +1071,7 @@ async def test_test_cache_connection_does_not_log_plaintext_credentials(monkeypa existing.cache_settings = {"type": "redis", "host": "h", "port": "6379", "password": "realredispw"} mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db proxy_config = MagicMock() proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) @@ -1095,6 +1110,7 @@ async def test_test_cache_connection_does_not_replay_saved_password_to_new_host( existing.cache_settings = {"type": "redis", "host": "real-redis", "port": "6379", "password": "realredispw"} mock_prisma = MagicMock() mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + mock_prisma.replica_db = mock_prisma.db proxy_config = MagicMock() proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index 69013408962..6eb6a543554 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -387,6 +387,7 @@ class TestTeamAdminCanInviteUser: teams = [make_team(tid, tid in user_is_admin_in) for tid in admin_teams] mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=teams) + mock_prisma.replica_db = mock_prisma.db result = await _team_admin_can_invite_user( user_api_key_dict=mock_auth, @@ -1072,6 +1073,7 @@ class TestTeamAdminCanInviteUserQuery: find_many = AsyncMock(return_value=[make_team("t1"), make_team("t2")]) mock_prisma.db.litellm_teamtable.find_many = find_many + mock_prisma.replica_db = mock_prisma.db await _team_admin_can_invite_user( user_api_key_dict=mock_auth, @@ -1281,7 +1283,8 @@ async def test_router_weights_validate_current_deployment_scope( }]) rows = [SimpleNamespace(model_id="id", model_name=stored_name, model_info=info)] if stored_name else [] table = SimpleNamespace(find_many=AsyncMock(return_value=rows)) - db = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table)) + tables = SimpleNamespace(litellm_proxymodeltable=table) + db = SimpleNamespace(db=tables, replica_db=tables) validation = validate_router_settings_weights( {"weights": {"group": {"id": 1}}}, team_id="team", prisma_client=db, llm_router=router, ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py index 49b0ed1b28a..6ebf4c46500 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_config_override_endpoints.py @@ -38,6 +38,7 @@ def _make_mock_db(): mock.delete = AsyncMock(return_value=None) prisma = MagicMock() prisma.db.litellm_configoverrides = mock + prisma.replica_db = prisma.db return prisma, mock diff --git a/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py index 7c6e8154107..bde5d48bb7f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_coordination_redis_endpoints.py @@ -52,6 +52,7 @@ def _prisma_with_general_settings(general_settings: dict | None) -> MagicMock: mock_prisma = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=row) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.upsert = AsyncMock() return mock_prisma @@ -267,6 +268,7 @@ async def test_update_rejects_settings_without_a_connection_target(monkeypatch): assert exc_info.value.status_code == 400 mock_prisma.db.litellm_config.upsert.assert_not_called() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.asyncio @@ -650,6 +652,7 @@ async def test_update_refuses_a_config_owned_coordination_redis_block(monkeypatc assert refused.value.status_code == 400 assert refused.value.detail["keys"] == ["coordination_redis"] mock_prisma.db.litellm_config.upsert.assert_not_called() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_budget.py b/tests/test_litellm/proxy/management_endpoints/test_customer_budget.py index 0beca0c15e8..a6d58ad8bc5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_budget.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_budget.py @@ -69,6 +69,7 @@ async def test_update_customer_with_budget_id( mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=mock_existing_customer ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_updated_user = MagicMock() mock_updated_user.model_dump.return_value = { @@ -125,6 +126,7 @@ async def test_update_customer_creates_budget_with_proper_relations( mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=mock_existing_customer ) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock budget creation mock_created_budget = MagicMock() @@ -183,6 +185,7 @@ async def test_update_customer_creates_budget_with_required_fields( mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=mock_existing_customer ) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock budget creation mock_created_budget = MagicMock() @@ -248,6 +251,7 @@ async def test_update_customer_budget_creation_with_fallback_admin( mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=mock_existing_customer ) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock budget creation mock_created_budget = MagicMock() @@ -305,6 +309,7 @@ async def test_update_customer_with_budget_id_and_creation_fields( mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=mock_existing_customer ) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock budget creation mock_created_budget = MagicMock() diff --git a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py index 77e52f30bb7..cc36549002b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_customer_endpoints.py @@ -67,6 +67,7 @@ def test_update_customer_success(mock_prisma_client, mock_user_api_key_auth): # Mock the find_first response mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=mock_end_user) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock the update response mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=updated_mock_end_user) @@ -88,6 +89,7 @@ def test_update_customer_unblock(mock_prisma_client, mock_user_api_key_auth): updated_mock_end_user = LiteLLM_EndUserTable(user_id="test-user-1", blocked=False) mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=mock_end_user) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=updated_mock_end_user) response = client.post( @@ -113,6 +115,7 @@ def test_update_customer_keeps_blocked_when_omitted(mock_prisma_client, mock_use updated_mock_end_user = LiteLLM_EndUserTable(user_id="test-user-1", blocked=True) mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=mock_end_user) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=updated_mock_end_user) response = client.post( @@ -133,6 +136,7 @@ def test_update_customer_not_found(mock_prisma_client, mock_user_api_key_auth): """ # Mock the database response to return None (user not found) mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db # Test data test_data = {"user_id": "non-existent-user", "alias": "Test User"} @@ -160,6 +164,7 @@ def test_info_customer_not_found(mock_prisma_client, mock_user_api_key_auth): """ # Mock the database response to return None (user not found) mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db # Make the request response = client.get( @@ -183,6 +188,7 @@ def test_delete_customer_not_found(mock_prisma_client, mock_user_api_key_auth): """ # Mock the database response to return empty list (no users found) mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.replica_db = mock_prisma_client.db # Test data test_data = {"user_ids": ["non-existent-user-1", "non-existent-user-2"]} @@ -225,6 +231,7 @@ def test_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth): # Test /customer/info - not found error mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db response = client.get( "/customer/info?end_user_id=non-existent", headers={"Authorization": "Bearer test-key"}, @@ -286,6 +293,7 @@ def test_customer_endpoints_error_schema_consistency(mock_prisma_client, mock_us # Scenario 1: GET /end_user/info with non-existent user # Should return 404 with proper error schema mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db response1 = client.get( "/end_user/info?end_user_id=fake-test-end-user-michaels-local-testng", @@ -386,6 +394,7 @@ def test_update_customer_response_preserves_budget_id(mock_prisma_client, mock_u existing = LiteLLM_EndUserTable(user_id="cust-1", blocked=False) updated = LiteLLM_EndUserTable(user_id="cust-1", blocked=False, budget_id="budget-123") mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=existing) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=updated) response = client.post( @@ -444,6 +453,7 @@ def test_update_customer_budget_omission_and_null_preserve_existing_budget( return LiteLLM_BudgetTable(budget_id="budget-1", max_budget=budget_state.max_budget) mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=end_user_row()) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.update = AsyncMock(side_effect=update_budget) mock_prisma_client.db.litellm_endusertable.update = AsyncMock(side_effect=lambda **_: response_row()) @@ -488,6 +498,7 @@ def test_update_customer_response_keeps_nested_budget_server_fields(mock_prisma_ }, } mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=existing) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=raw_row) response = client.post( @@ -512,6 +523,7 @@ def test_block_customer_success_serializes_through_response_model(mock_prisma_cl """ blocked_row = LiteLLM_EndUserTable(user_id="blocked-1", blocked=True) mock_prisma_client.db.litellm_endusertable.upsert = AsyncMock(return_value=blocked_row) + mock_prisma_client.replica_db = mock_prisma_client.db response = client.post( "/customer/block", @@ -535,6 +547,7 @@ def test_delete_customer_success_serializes_through_response_model(mock_prisma_c LiteLLM_EndUserTable(user_id="u2", blocked=False), ] mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=existing) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.delete_many = AsyncMock(return_value=2) response = client.post( @@ -560,6 +573,7 @@ async def test_get_customer_daily_activity_admin_param_passing(monkeypatch): mock_prisma_client = AsyncMock() mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocked_response = MagicMock(name="SpendAnalyticsPaginatedResponse") @@ -612,6 +626,7 @@ async def test_get_customer_daily_activity_with_end_user_aliases(monkeypatch): mock_end_user2.alias = "Customer Two" mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[mock_end_user1, mock_end_user2]) + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocked_response = MagicMock(name="SpendAnalyticsPaginatedResponse") @@ -838,6 +853,7 @@ def _row(dump: dict) -> MagicMock: def test_char_info_body(mock_prisma_client, mock_user_api_key_auth): mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=_row(_FULL_DB_ROW)) + mock_prisma_client.replica_db = mock_prisma_client.db response = client.get("/customer/info?end_user_id=c1", headers={"Authorization": "Bearer k"}) assert response.status_code == 200 assert response.json() == _EXPECTED_CUSTOMER @@ -845,6 +861,7 @@ def test_char_info_body(mock_prisma_client, mock_user_api_key_auth): def test_char_list_body(mock_prisma_client, mock_user_api_key_auth): mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[_row(_FULL_DB_ROW)]) + mock_prisma_client.replica_db = mock_prisma_client.db response = client.get("/customer/list", headers={"Authorization": "Bearer k"}) assert response.status_code == 200 assert response.json() == [_EXPECTED_CUSTOMER] @@ -852,6 +869,7 @@ def test_char_list_body(mock_prisma_client, mock_user_api_key_auth): def test_char_new_body(mock_prisma_client, mock_user_api_key_auth): mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW)) + mock_prisma_client.replica_db = mock_prisma_client.db response = client.post("/customer/new", json={"user_id": "c1"}, headers={"Authorization": "Bearer k"}) assert response.status_code == 200 assert response.json() == _EXPECTED_CUSTOMER @@ -864,6 +882,7 @@ def test_customer_new_rejects_a_duration_that_never_advances( """A zero-length window resets to "now", leaving the customer's budget row permanently due for the reset job to re-read every tick.""" mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW)) + mock_prisma_client.replica_db = mock_prisma_client.db response = client.post( "/customer/new", @@ -878,6 +897,7 @@ def test_customer_new_rejects_a_duration_that_never_advances( def test_customer_new_accepts_a_normal_duration(mock_prisma_client, mock_user_api_key_auth): mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW)) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.create = AsyncMock( return_value=_row({"budget_id": "b1", "max_budget": 10.0}) ) @@ -895,6 +915,7 @@ def test_char_update_body(mock_prisma_client, mock_user_api_key_auth): mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=_row({"user_id": "c1", "blocked": False}) ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=_row(_FULL_DB_ROW)) response = client.post( "/customer/update", @@ -912,6 +933,7 @@ def test_char_delete_body(mock_prisma_client, mock_user_api_key_auth): LiteLLM_EndUserTable(user_id="c2", blocked=False), ] ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.delete_many = AsyncMock(return_value=2) response = client.post( "/customer/delete", @@ -963,6 +985,7 @@ def test_customer_new_invalidates_end_user_and_registry_caches(mock_prisma_clien budget or block goes unenforced until the TTL expires. """ mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW)) + mock_prisma_client.replica_db = mock_prisma_client.db with _end_user_cache_doubles() as (recording_cache, mock_publish): response = client.post( @@ -981,6 +1004,7 @@ def test_customer_update_invalidates_end_user_and_registry_caches(mock_prisma_cl mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock( return_value=_row({"user_id": "c1", "blocked": False}) ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=_row(_FULL_DB_ROW)) with _end_user_cache_doubles() as (recording_cache, mock_publish): @@ -1000,6 +1024,7 @@ def test_customer_block_invalidates_end_user_and_registry_caches(mock_prisma_cli mock_prisma_client.db.litellm_endusertable.upsert = AsyncMock( return_value=LiteLLM_EndUserTable(user_id="c1", blocked=True) ) + mock_prisma_client.replica_db = mock_prisma_client.db with _end_user_cache_doubles() as (recording_cache, mock_publish): response = client.post( @@ -1029,6 +1054,7 @@ def test_customer_delete_invalidates_end_user_and_registry_caches(mock_prisma_cl LiteLLM_EndUserTable(user_id="c2", blocked=False), ] ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_endusertable.delete_many = AsyncMock(return_value=2) with _end_user_cache_doubles() as (recording_cache, mock_publish): diff --git a/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py b/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py index e33945df7dc..888b70b9428 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py +++ b/tests/test_litellm/proxy/management_endpoints/test_delete_verification_tokens_failed.py @@ -69,6 +69,7 @@ def _mock_prisma(keys, deleted_tokens): """Return a minimal mock prisma_client for a given set of found keys and deleted tokens.""" mock = AsyncMock() mock.db.litellm_verificationtoken.find_many = AsyncMock(return_value=keys) + mock.replica_db = mock.db mock.delete_data = AsyncMock(return_value=deleted_tokens) mock.db.litellm_deletedverificationtoken.create_many = AsyncMock() return mock @@ -176,6 +177,7 @@ async def test_delete_tokens_non_admin_token_not_in_db_returns_failed_tokens( mock_prisma = AsyncMock() # DB find_many returns only key1 — token-2 is not found mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[key1]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.delete_data = AsyncMock(return_value=["hashed-token-1"]) mock_prisma.db.litellm_deletedverificationtoken.create_many = AsyncMock() diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 557e753a76f..040be87e356 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -120,6 +120,7 @@ def setup_mock_prisma_client( ): """Helper to set up a mock prisma client with proper async behavior""" mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable = AsyncMock() mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=team_records) mock_prisma_client.db.litellm_mcpservertable = AsyncMock() @@ -1334,6 +1335,7 @@ class TestListMCPServers: mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=raw_prisma_model) + mock_prisma_client.replica_db = mock_prisma_client.db mock_health_result = generate_mock_mcp_server_db_record(server_id="env-server", alias="Env Server") mock_health_result.status = "healthy" @@ -3845,6 +3847,7 @@ class TestUpdateMCPServer: # Mock dependencies mock_prisma_client = MagicMock() mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_mcpservertable = AsyncMock() mock_prisma_client.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_server) mock_prisma_client.db.litellm_mcpservertable.update = AsyncMock(return_value=updated_server) @@ -4973,6 +4976,7 @@ def _make_prisma_client(): """Return a minimal mock PrismaClient accepted by get_prisma_client_or_throw.""" client = MagicMock() client.db = MagicMock() + client.replica_db = client.db return client @@ -7916,6 +7920,7 @@ async def test_config_server_edit_preserves_api_contract_without_creating_rows(r original = server.model_dump() prisma = MagicMock() prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=None) + prisma.replica_db = prisma.db prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=None) with ( patch.object(mgmt_endpoints, "global_mcp_server_manager", manager), diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index 3c6afa86c45..8091736861d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -59,6 +59,7 @@ async def test_organization_update_object_permissions_existing_permission(monkey } mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_object_permission) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock upsert operation updated_permission = MagicMock() @@ -104,6 +105,7 @@ async def test_get_organization_daily_activity_admin_param_passing(monkeypatch): # Mock prisma client mock_prisma_client = AsyncMock() mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Admin view -> skip membership restriction @@ -165,6 +167,7 @@ async def test_get_organization_daily_activity_non_admin_defaults_to_admin_orgs( # Mock prisma client and memberships mock_prisma_client = AsyncMock() mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[]) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_organizationmembership.find_many = AsyncMock( return_value=[ SimpleNamespace(organization_id="orgA", user_role=LitellmUserRoles.ORG_ADMIN.value), @@ -222,6 +225,7 @@ async def test_get_organization_daily_activity_non_admin_unauthorized_org_raises mock_prisma_client.db.litellm_organizationmembership.find_many = AsyncMock( return_value=[SimpleNamespace(organization_id="orgA", user_role=LitellmUserRoles.ORG_ADMIN.value)] ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -287,6 +291,7 @@ async def test_organization_update_object_permissions_no_existing_permission( # Mock find_unique to return None (no existing permission) mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock upsert to create new record new_permission = MagicMock() @@ -352,6 +357,7 @@ async def test_organization_update_object_permissions_missing_permission_record( # Mock find_unique to return None (permission record not found) mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock upsert to create new record new_permission = MagicMock() @@ -413,6 +419,7 @@ async def test_list_organization_filter_by_org_id(monkeypatch): # Mock find_many to return filtered results mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[mock_org1]) + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -475,6 +482,7 @@ async def test_list_organization_filter_by_org_alias(monkeypatch): # Mock find_many to return filtered results mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[mock_org1, mock_org2]) + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -570,6 +578,7 @@ def patched_org_prisma(): patch("litellm.proxy.proxy_server.proxy_logging_obj"), ): mock_prisma.db.litellm_organizationtable.find_unique = AsyncMock(return_value=victim_row) + mock_prisma.replica_db = mock_prisma.db yield mock_prisma @@ -662,7 +671,7 @@ async def test_organization_member_add_budget_omission_and_null_leave_budget_uns litellm_usertable=SimpleNamespace(find_unique=AsyncMock(return_value=user)), litellm_organizationmembership=SimpleNamespace(create=create_membership), ) - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, replica_db=mock_db) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) monkeypatch.setattr( "litellm.proxy.management_endpoints.organization_endpoints._verify_org_access", @@ -737,7 +746,7 @@ async def test_organization_member_update_budget_omission_and_null_preserve_exis find_unique=AsyncMock(return_value=SimpleNamespace(user_role="internal_user")) ), ) - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, replica_db=mock_db) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) monkeypatch.setattr(organization_endpoints, "update_budget", update_budget) monkeypatch.setattr( @@ -833,6 +842,7 @@ async def _run_update_organization_v2( existing_org.metadata = existing_metadata mock_prisma_client.db.litellm_organizationtable.find_unique = AsyncMock(return_value=existing_org) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_organizationtable.update = AsyncMock(return_value=MagicMock()) mock_prisma_client.db.litellm_budgettable.update = AsyncMock() mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( @@ -905,6 +915,7 @@ async def test_v2_update_untouched_fields_not_written(monkeypatch): ) prisma.db.litellm_budgettable.update.assert_not_awaited() + prisma.replica_db = prisma.db write_data = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"] assert write_data["organization_alias"] == "renamed" assert "metadata" not in write_data @@ -982,6 +993,7 @@ async def test_v2_rejects_negative_integer_limits(monkeypatch: pytest.MonkeyPatc assert exc.value.status_code == 422 assert field in str(exc.value.detail) prisma_mock.db.tx.assert_not_called() + prisma_mock.replica_db = prisma_mock.db prisma_mock.db.litellm_budgettable.update.assert_not_awaited() prisma_mock.db.litellm_organizationtable.update.assert_not_awaited() @@ -1006,6 +1018,7 @@ async def test_v2_rejects_unparseable_budget_duration(monkeypatch: pytest.Monkey assert exc.value.status_code == 422 assert "budget_duration" in str(exc.value.detail) prisma_mock.db.tx.assert_not_called() + prisma_mock.replica_db = prisma_mock.db prisma_mock.db.litellm_budgettable.update.assert_not_awaited() prisma_mock.db.litellm_organizationtable.update.assert_not_awaited() @@ -1034,6 +1047,7 @@ async def test_v2_rejects_caller_without_org_access(monkeypatch): ) assert exc.value.status_code == 403 mock_prisma_client.db.litellm_organizationtable.update.assert_not_awaited() + mock_prisma_client.replica_db = mock_prisma_client.db @pytest.mark.asyncio @@ -1075,6 +1089,7 @@ async def test_v2_object_permission_upsert_runs_inside_transaction(monkeypatch): prisma.tx.litellm_objectpermissiontable.upsert.assert_awaited_once() prisma.db.litellm_objectpermissiontable.upsert.assert_not_awaited() + prisma.replica_db = prisma.db upsert = prisma.tx.litellm_objectpermissiontable.upsert.await_args.kwargs linked_id = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"]["object_permission_id"] @@ -1097,6 +1112,7 @@ async def test_v2_clears_object_permission_when_sent_null(monkeypatch): prisma.tx.litellm_objectpermissiontable.upsert.assert_not_awaited() prisma.db.litellm_objectpermissiontable.find_unique.assert_not_awaited() + prisma.replica_db = prisma.db write_data = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"] assert write_data["object_permission_id"] is None @@ -1122,6 +1138,7 @@ async def test_v2_rejects_empty_object_permission(monkeypatch): assert exc.value.status_code == 422 assert "object_permission" in str(exc.value.detail) mock_prisma_client.db.litellm_organizationtable.update.assert_not_awaited() + mock_prisma_client.replica_db = mock_prisma_client.db @pytest.mark.asyncio @@ -1135,6 +1152,7 @@ async def test_v2_writes_budget_and_org_in_one_transaction(monkeypatch): ) prisma.db.tx.assert_called_once() + prisma.replica_db = prisma.db prisma.db.litellm_budgettable.update.assert_awaited_once() prisma.db.litellm_organizationtable.update.assert_awaited_once() @@ -1174,6 +1192,7 @@ async def _run_legacy_update_organization( existing_org.budget_id = existing_budget_id existing_org.metadata = {} mock_prisma_client.db.litellm_organizationtable.find_unique = AsyncMock(return_value=existing_org) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_organizationtable.update = AsyncMock(return_value=MagicMock()) mock_prisma_client.db.litellm_budgettable.update = AsyncMock() @@ -1216,6 +1235,7 @@ async def test_legacy_update_without_budget_fields_skips_budget_write(monkeypatc ) prisma.db.litellm_budgettable.update.assert_not_awaited() + prisma.replica_db = prisma.db assert prisma.db.litellm_organizationtable.update.await_args.kwargs["data"]["organization_alias"] == "renamed" @@ -1267,6 +1287,7 @@ async def test_get_organization_daily_activity_non_admin_without_org_admin_role_ mock_prisma_client = AsyncMock() org_table_find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_organizationtable.find_many = org_table_find_many + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_organizationmembership.find_many = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -1308,6 +1329,7 @@ async def test_find_member_if_email_missing_row_raises_documented_400(): prisma_client = AsyncMock() prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db with pytest.raises(HTTPException) as exc_info: await find_member_if_email("missing@example.com", prisma_client) @@ -1337,6 +1359,7 @@ async def test_new_organization_rejects_shared_alias_tool_permission_key(): MagicMock(server_id="wiki-b-id", alias="wiki", server_name="wiki_b"), ] ) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_objectpermissiontable.create = AsyncMock() data = NewOrganizationRequest( organization_alias="org", @@ -1367,6 +1390,7 @@ async def test_new_organization_temp_budget_fields_go_to_budget_row_not_metadata prisma_client = MagicMock() prisma_client.jsonify_object = MagicMock(side_effect=lambda data: PrismaClient.jsonify_object(prisma_client, data)) prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_budgettable.create = AsyncMock(return_value=MagicMock(budget_id="budget-1")) prisma_client.db.litellm_organizationtable.create = AsyncMock(return_value={"organization_id": "org-1"}) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) @@ -1523,12 +1547,14 @@ async def test_delete_organization_evicts_the_cache_of_the_keys_it_deletes(monke prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[SimpleNamespace(token="hashed-org-key")] ) + prisma_client.replica_db = prisma_client.db async def cascading_delete_many(where): jwt_table.cascade(("hashed-org-key",)) return 1 prisma_client.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=cascading_delete_many) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping = jwt_table prisma_client.db.litellm_organizationtable.delete = AsyncMock(return_value=MagicMock()) diff --git a/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py index adff4eda47c..57570e492e9 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py @@ -35,6 +35,7 @@ def _make_user_row(password: str | None) -> MagicMock: def _make_prisma(user: MagicMock | None) -> MagicMock: prisma = MagicMock() prisma.db.litellm_usertable.find_first = AsyncMock(return_value=user) + prisma.replica_db = prisma.db prisma.db.litellm_usertable.update = AsyncMock(return_value=user) return prisma @@ -130,6 +131,7 @@ async def test_change_password_rejects_wrong_current_password(): assert exc_info.value.status_code == 400 assert "Current password is incorrect" in exc_info.value.detail["error"] prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -156,6 +158,7 @@ async def test_change_password_rejects_unchanged_password(): assert exc_info.value.status_code == 400 assert "must be different from the current password" in exc_info.value.detail["error"] prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -192,6 +195,7 @@ async def test_change_password_rejects_non_password_login_session(caller: UserAP assert exc_info.value.status_code == 403 assert "logging in with a password" in exc_info.value.detail["error"] prisma.db.litellm_usertable.find_first.assert_not_called() + prisma.replica_db = prisma.db prisma.db.litellm_usertable.update.assert_not_called() @@ -218,6 +222,7 @@ async def test_change_password_rejects_session_without_user(): assert exc_info.value.status_code == 400 prisma.db.litellm_usertable.find_first.assert_not_called() + prisma.replica_db = prisma.db prisma.db.litellm_usertable.update.assert_not_called() @@ -246,6 +251,7 @@ async def test_change_password_rejects_account_without_password(): assert exc_info.value.status_code == 400 assert "no password set" in exc_info.value.detail["error"] prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -274,6 +280,7 @@ async def test_change_password_enforces_min_length(): assert exc_info.value.param == "password" assert "at least 12 characters" in exc_info.value.message prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -304,6 +311,7 @@ async def test_change_password_rejects_breached_password(): assert exc_info.value.param == "password" assert "data breaches" in exc_info.value.message prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -335,6 +343,7 @@ async def test_change_password_verifies_current_password_before_hibp_lookup(): assert "Current password is incorrect" in exc_info.value.detail["error"] assert hibp_calls == [] prisma.db.litellm_usertable.update.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py b/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py index ce08030739d..1f92202ddb6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py +++ b/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py @@ -30,6 +30,7 @@ def _make_prisma_with_team(team_id: str, admins: list, members_with_roles: tuple prisma = MagicMock() team_row = LiteLLM_TeamTable(team_id=team_id, admins=admins, members_with_roles=list(members_with_roles)) prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + prisma.replica_db = prisma.db return prisma @@ -57,6 +58,7 @@ async def test_project_perm_check_uses_current_team_not_caller_supplied(): ) assert has_perm is False prisma.db.litellm_teamtable.find_unique.assert_awaited_once() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -130,6 +132,7 @@ async def test_project_perm_check_denies_team_admin_unless_projects_permission_c ) assert has_perm is False prisma.db.litellm_teamtable.find_unique.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -151,6 +154,7 @@ async def test_project_perm_check_require_admin_denies_team_admin_even_when_conf ) assert has_perm is False prisma.db.litellm_teamtable.find_unique.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -172,6 +176,7 @@ async def test_project_perm_check_uses_injected_team_object_for_reassignment_tar ) assert has_perm is False prisma.db.litellm_teamtable.find_unique.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -195,6 +200,7 @@ async def test_project_perm_check_proxy_admin_always_allowed(): assert has_perm is True # Admin shortcut should not even hit the DB. prisma.db.litellm_teamtable.find_unique.assert_not_called() + prisma.replica_db = prisma.db # --------------------------------------------------------------------------- @@ -209,6 +215,7 @@ def _make_prisma_with_user_orgs(user_id: str, org_ids: list): MagicMock(organization_id=org_id) for org_id in org_ids ] prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + prisma.replica_db = prisma.db return prisma @@ -282,6 +289,7 @@ async def test_assign_key_org_blocks_caller_with_no_memberships(): user_row = MagicMock() user_row.organization_memberships = None prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + prisma.replica_db = prisma.db caller = UserAPIKeyAuth( user_id="alice", diff --git a/tests/test_litellm/proxy/management_endpoints/test_session_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_session_endpoints.py index d5960a88937..6f2c8e12f7b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_session_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_session_endpoints.py @@ -29,6 +29,7 @@ def _make_prisma( ) -> MagicMock: prisma = MagicMock() table = prisma.db.litellm_verificationtoken + prisma.replica_db = prisma.db table.find_unique = AsyncMock(return_value=find_unique_row) table.find_many = AsyncMock(return_value=find_many_rows or []) table.delete_many = AsyncMock(return_value=1) @@ -133,6 +134,7 @@ async def test_session_logout_refuses_non_ui_session_key(): assert exc_info.value.status_code == 403 prisma.db.litellm_verificationtoken.delete_many.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -157,6 +159,7 @@ async def test_session_logout_is_idempotent_when_row_already_gone(): assert response.message == "Session already revoked." prisma.db.litellm_verificationtoken.delete_many.assert_not_called() + prisma.replica_db = prisma.db # The cache entry may outlive the row; evict regardless. evict_mock.assert_awaited_once() @@ -260,6 +263,7 @@ async def test_revoke_ui_session_keys_noop_when_no_sessions(): assert revoked == 0 prisma.db.litellm_verificationtoken.delete_many.assert_not_called() + prisma.replica_db = prisma.db @pytest.mark.asyncio @@ -268,6 +272,7 @@ async def test_revoke_ui_session_keys_failure_is_swallowed(): failure must not fail the caller's request.""" prisma = _make_prisma(find_many_rows=[_session_row(token="t1")]) prisma.db.litellm_verificationtoken.delete_many = AsyncMock(side_effect=RuntimeError("db down")) + prisma.replica_db = prisma.db p1, p2, p3 = _patched_globals(prisma) with ( diff --git a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py index 3cfdd345a45..1062b75940e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_tag_management_endpoints.py @@ -87,6 +87,7 @@ async def test_create_and_get_tag(): # Setup prisma mocks mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock find_unique to return None (tag doesn't exist) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=None) @@ -179,6 +180,7 @@ async def test_update_tag(): # Setup prisma mocks mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock existing tag existing_tag = Mock() @@ -247,7 +249,7 @@ async def test_new_tag_persists_a_budget(): created_by="admin", ) mock_db = Mock() - mock_prisma = SimpleNamespace(db=mock_db, jsonify_object=lambda data: dict(data)) + mock_prisma = SimpleNamespace(db=mock_db, replica_db=mock_db, jsonify_object=lambda data: dict(data)) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=None) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) @@ -316,7 +318,7 @@ async def test_update_tag_explicit_null_preserves_general_budget_fields(field): created_by="admin", ) mock_db = Mock() - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, replica_db=mock_db) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) @@ -370,7 +372,7 @@ async def test_update_tag_explicit_null_clears_budget_duration(): created_by="admin", ) mock_db = Mock() - mock_prisma = SimpleNamespace(db=mock_db) + mock_prisma = SimpleNamespace(db=mock_db, replica_db=mock_db) mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=existing_tag) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_db.litellm_tagtable.update = AsyncMock(return_value=updated_tag) @@ -420,6 +422,7 @@ async def test_delete_tag(): # Setup prisma mocks mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock existing tag existing_tag = Mock() @@ -516,6 +519,7 @@ async def test_new_tag_invalidates_tag_and_registry_caches(): ): mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_db.litellm_tagtable.find_unique = AsyncMock(return_value=None) mock_db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_get_deployments.return_value = [] @@ -568,6 +572,7 @@ async def test_update_tag_invalidates_only_the_tag_cache(): ): mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db existing_tag = Mock() existing_tag.tag_name = "cache-tag" @@ -619,6 +624,7 @@ async def test_delete_tag_invalidates_tag_and_registry_caches(): ): mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db existing_tag = Mock() existing_tag.tag_name = "cache-tag" @@ -659,6 +665,7 @@ async def test_list_tags_with_dynamic_tags(): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Setup stored tags stored_tag = Mock() @@ -740,6 +747,7 @@ async def test_list_tags_no_dynamic_tags(): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db stored_tag = Mock() stored_tag.tag_name = "stored-tag" @@ -790,6 +798,7 @@ async def test_internal_user_list_tags_only_returns_tags_used_by_their_keys(): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db owned_key_record = Mock() owned_key_record.token = "owned-key" @@ -881,6 +890,7 @@ async def test_internal_user_list_tags_does_not_500_on_unsupported_prisma_kwarg( with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db key_record = Mock() key_record.token = "new-user-key" @@ -923,6 +933,7 @@ async def test_list_tags_with_date_range_filters_dynamic_tags(): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[]) group_by_mock = AsyncMock(return_value=[]) mock_db.litellm_dailytagspend.group_by = group_by_mock @@ -969,6 +980,7 @@ async def test_internal_user_tag_daily_activity_is_scoped_to_their_keys(): ): mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db owned_key_record = Mock() owned_key_record.token = "owned-key" @@ -1015,6 +1027,7 @@ async def test_internal_user_tag_daily_activity_rejects_unowned_api_key_filter() ): mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db owned_key_record = Mock() owned_key_record.token = "owned-key" @@ -1061,6 +1074,7 @@ async def test_internal_user_tag_daily_activity_scopes_to_current_key_without_us ): mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db fake_token_table = FakeVerificationTokenTable([]) mock_db.litellm_verificationtoken = fake_token_table mock_get_daily_activity.return_value = "daily-activity-response" @@ -1105,6 +1119,7 @@ async def test_internal_user_tag_daily_activity_without_any_scoped_keys_returns_ ): mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db fake_token_table = FakeVerificationTokenTable([]) mock_db.litellm_verificationtoken = fake_token_table @@ -1165,6 +1180,7 @@ async def test_list_tags_without_date_range_omits_date_filter(): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[]) group_by_mock = AsyncMock(return_value=[]) mock_db.litellm_dailytagspend.group_by = group_by_mock @@ -1206,6 +1222,7 @@ async def test_list_tags_rejects_invalid_date_range(query, expected_detail_fragm with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_db.litellm_tagtable.find_many = AsyncMock(return_value=[]) mock_db.litellm_dailytagspend.group_by = AsyncMock(return_value=[]) @@ -1325,6 +1342,7 @@ async def test_add_tag_to_deployment_preserves_encrypted_fields(): # Setup prisma mocks mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock the database model with encrypted fields db_model = Mock() @@ -1391,6 +1409,7 @@ async def test_add_tag_to_deployment_with_string_params(): # Setup prisma mocks mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock the database model with litellm_params as string db_model = Mock() @@ -1444,6 +1463,7 @@ async def test_add_tag_to_deployment_no_duplicate_tags(): # Setup prisma mocks mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock the database model with existing tags db_model = Mock() @@ -1496,6 +1516,7 @@ async def test_add_tag_to_deployment_model_not_found(): # Setup prisma mocks mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock find_unique to return None (model not found) mock_db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py index b6eebcb2ef3..a82415ba3ef 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py @@ -57,6 +57,7 @@ def _patch_prisma(existing_team: MagicMock): updated_row = MagicMock() updated_row.team_id = existing_team.team_id mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=updated_row) + mock_prisma.replica_db = mock_prisma.db return mock_prisma @@ -106,6 +107,7 @@ def patched_prisma(): ): mock_client.get_data = AsyncMock(return_value=_team_row()) mock_client.db.litellm_teamtable.update = AsyncMock() + mock_client.replica_db = mock_client.db yield mock_client @@ -128,6 +130,7 @@ async def test_add_team_callbacks_rejects_unauthorized_caller(patched_prisma, un ) assert exc.value.status_code == 403 patched_prisma.db.litellm_teamtable.update.assert_not_called() + patched_prisma.replica_db = patched_prisma.db @pytest.mark.asyncio @@ -140,6 +143,7 @@ async def test_disable_team_logging_rejects_unauthorized_caller(patched_prisma, ) assert exc.value.status_code == 403 patched_prisma.db.litellm_teamtable.update.assert_not_called() + patched_prisma.replica_db = patched_prisma.db @pytest.mark.asyncio @@ -170,6 +174,7 @@ async def test_proxy_admin_can_add_team_callbacks(patched_prisma): user_api_key_dict=_admin_auth(), ) patched_prisma.db.litellm_teamtable.update.assert_awaited_once() + patched_prisma.replica_db = patched_prisma.db @pytest.mark.asyncio @@ -196,6 +201,7 @@ async def test_team_admin_of_target_team_can_add_callbacks(patched_prisma): user_api_key_dict=team_admin, ) patched_prisma.db.litellm_teamtable.update.assert_awaited_once() + patched_prisma.replica_db = patched_prisma.db @pytest.mark.asyncio @@ -971,6 +977,7 @@ async def test_delete_team_callback_rejects_unauthorized_caller(patched_prisma, ) assert exc.value.status_code == 403 patched_prisma.db.litellm_teamtable.update.assert_not_called() + patched_prisma.replica_db = patched_prisma.db @pytest.mark.asyncio @@ -1109,6 +1116,7 @@ async def test_delete_team_callback_404s_for_unregistered_callback(): assert exc.value.status_code == 404 assert exc.value.detail == {"error": "callback_name = gcs is not registered for team_id = team-1."} mock_prisma.db.litellm_teamtable.update.assert_not_called() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.asyncio @@ -1138,6 +1146,7 @@ async def test_delete_team_callback_404s_when_team_has_no_logging_slot(): assert exc.value.status_code == 404 mock_prisma.db.litellm_teamtable.update.assert_not_called() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.asyncio @@ -1145,6 +1154,7 @@ async def test_delete_team_callback_404s_for_unknown_team(): mock_prisma = MagicMock() mock_prisma.get_data = AsyncMock(return_value=None) mock_prisma.db.litellm_teamtable.update = AsyncMock() + mock_prisma.replica_db = mock_prisma.db with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with pytest.raises(HTTPException) as exc: @@ -1171,6 +1181,7 @@ async def test_add_team_callbacks_rejects_team_deleted_before_write(): """ mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata={})) mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db data = AddTeamCallback( callback_name="langfuse", @@ -1488,6 +1499,7 @@ async def test_unknown_team_is_indistinguishable_from_no_access(call_handler, un ): # test-quality-ok: the handler imports prisma_client from proxy_server at call time, so there is no seam to inject through mock_client.get_data = AsyncMock(return_value=_team_row()) mock_client.db.litellm_teamtable.update = AsyncMock() + mock_client.replica_db = mock_client.db with patch( # test-quality-ok: _verify_team_access calls this module-level helper directly, so there is no seam to inject through "litellm.proxy.management_endpoints.team_endpoints._is_user_org_admin_for_team", new_callable=AsyncMock, @@ -1663,6 +1675,7 @@ async def test_a_second_entry_may_not_flip_the_span_scope(patched_prisma, caller assert exc.value.status_code == 400 assert "langfuse_span_scope" in str(exc.value.detail) and "'full'" in str(exc.value.detail) patched_prisma.db.litellm_teamtable.update.assert_not_called() + patched_prisma.replica_db = patched_prisma.db data.callback_vars["langfuse_span_scope"] = "full" await add_team_callbacks( @@ -1691,6 +1704,7 @@ async def test_add_team_callbacks_rejects_out_of_range_arize_sampling_rate(patch assert exc.value.status_code == 400 assert "arize_success_sampling_rate" in str(exc.value.detail) patched_prisma.db.litellm_teamtable.update.assert_not_called() + patched_prisma.replica_db = patched_prisma.db def test_add_team_callback_accepts_arize_sampling_rate_vars(): diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py b/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py index 17cb30dd07d..1aebb35c8f5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_default_params.py @@ -147,6 +147,7 @@ class TestNewTeamDefaultParamsApplied: ) mock_prisma.get_generic_data = AsyncMock(return_value=None) mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable = MagicMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -682,6 +683,7 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a, team_b]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -718,6 +720,7 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_has, team_missing]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -747,6 +750,7 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_many = AsyncMock(side_effect=[page1, page2]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -779,6 +783,7 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a, team_b]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -811,6 +816,7 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_has, team_missing]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.batch_ = MagicMock(return_value=mock_batcher) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -838,6 +844,7 @@ class TestBulkUpdateTeamMemberPermissions: mock_prisma = MagicMock() # Only team-a exists, team-b does not mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a]) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) data = BulkUpdateTeamMemberPermissionsRequest( @@ -914,6 +921,7 @@ class TestBulkUpdateTeamMemberPermissions: assert result["teams_updated"] == 0 mock_prisma.db.litellm_teamtable.find_many.assert_not_called() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.asyncio async def test_non_admin_gets_403(self, monkeypatch): diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_model_alias_merge.py b/tests/test_litellm/proxy/management_endpoints/test_team_model_alias_merge.py index 7cdf60f043e..05c8b5ae3cb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_model_alias_merge.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_model_alias_merge.py @@ -58,6 +58,7 @@ class TestTeamModelAddAtomicAppend: mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( return_value=existing_team ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.execute_raw = AsyncMock(return_value=None) mock_prisma.db.litellm_teamtable.update = AsyncMock( return_value=updated_team diff --git a/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py index 09d14cfe5df..45cda95705e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_tool_management_endpoints.py @@ -110,6 +110,7 @@ def _team_row(object_permission_id: Optional[str]) -> MagicMock: def _team_policy_prisma(team_table: FakeTeamTable) -> MagicMock: prisma = MagicMock() prisma.db.litellm_teamtable = team_table + prisma.replica_db = prisma.db prisma.db.litellm_objectpermissiontable.create = AsyncMock() prisma.db.litellm_objectpermissiontable.delete = AsyncMock() return prisma @@ -118,6 +119,7 @@ def _team_policy_prisma(team_table: FakeTeamTable) -> MagicMock: def _rollup_prisma(group_rows: list, daily_rows: list | None = None) -> MagicMock: prisma = MagicMock() prisma.db.query_raw = AsyncMock(return_value=[]) + prisma.replica_db = prisma.db prisma.db.litellm_spendlogs.find_many = AsyncMock(return_value=[]) prisma.db.litellm_spendlogtoolindex.find_many = AsyncMock(return_value=[]) prisma.db.litellm_dailytoolspend.group_by = AsyncMock(return_value=group_rows) @@ -343,6 +345,7 @@ class TestToolManagementEndpoints: resp = self.client.get("/v1/tool/spend?start_date=2026-07-01&end_date=2026-07-02") assert resp.status_code == 200 prisma.db.litellm_dailytoolspend.find_many.assert_not_awaited() + prisma.replica_db = prisma.db @patch("litellm.proxy.proxy_server.prisma_client", None) def test_tool_spend_no_db_returns_500(self): @@ -358,6 +361,7 @@ class TestToolManagementEndpoints: resp = self.client.get("/v1/tool/spend?start_date=2026-07-01&end_date=2026-07-02") assert resp.status_code == 200 prisma.db.query_raw.assert_not_awaited() + prisma.replica_db = prisma.db prisma.db.litellm_spendlogs.find_many.assert_not_awaited() prisma.db.litellm_spendlogtoolindex.find_many.assert_not_awaited() prisma.db.litellm_dailytoolspend.group_by.assert_awaited_once() @@ -409,6 +413,7 @@ class TestToolManagementEndpoints: assert resp.status_code == 400 assert "Invalid date format" in resp.json()["detail"] prisma.db.litellm_dailytoolspend.group_by.assert_not_awaited() + prisma.replica_db = prisma.db def test_tool_spend_non_admin_returns_403(self): from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth @@ -424,3 +429,4 @@ class TestToolManagementEndpoints: resp = client.get("/v1/tool/spend") assert resp.status_code == 403 prisma.db.litellm_dailytoolspend.group_by.assert_not_awaited() + prisma.replica_db = prisma.db diff --git a/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py index 0c3d5107cb6..20d2ebb5b4e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_workflow_management_endpoints.py @@ -104,6 +104,7 @@ def _make_tx(event_return=None, run_return=None, msg_return=None) -> MagicMock: def _make_prisma_client() -> MagicMock: client = MagicMock() client.db = MagicMock() + client.replica_db = client.db client.db.litellm_workflowrun = MagicMock() client.db.litellm_workflowevent = MagicMock() client.db.litellm_workflowmessage = MagicMock() @@ -185,6 +186,7 @@ class TestCreateWorkflowRun: @patch("litellm.proxy.proxy_server.prisma_client") def test_create_returns_run(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.create = AsyncMock(return_value=_make_run()) resp = self.client.post( @@ -215,6 +217,7 @@ class TestListWorkflowRuns: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_returns_runs(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_many = AsyncMock( return_value=[_make_run()] ) @@ -227,6 +230,7 @@ class TestListWorkflowRuns: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_filters_by_status(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[]) resp = self.client.get("/v1/workflows/runs?status=running") @@ -237,6 +241,7 @@ class TestListWorkflowRuns: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_filters_by_multiple_statuses(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[]) resp = self.client.get("/v1/workflows/runs?status=running,paused") @@ -257,6 +262,7 @@ class TestGetWorkflowRun: @patch("litellm.proxy.proxy_server.prisma_client") def test_get_existing_run(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -267,6 +273,7 @@ class TestGetWorkflowRun: @patch("litellm.proxy.proxy_server.prisma_client") def test_get_missing_run_returns_404(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None) resp = self.client.get("/v1/workflows/runs/nonexistent") @@ -285,6 +292,7 @@ class TestUpdateWorkflowRun: @patch("litellm.proxy.proxy_server.prisma_client") def test_update_status(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -300,6 +308,7 @@ class TestUpdateWorkflowRun: @patch("litellm.proxy.proxy_server.prisma_client") def test_update_no_fields_returns_400(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db resp = self.client.patch("/v1/workflows/runs/run-1", json={}) assert resp.status_code == 400 @@ -316,6 +325,7 @@ class TestAppendWorkflowEvent: @patch("litellm.proxy.proxy_server.prisma_client") def test_append_event_updates_run_status(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db # _require_run check self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() @@ -339,6 +349,7 @@ class TestAppendWorkflowEvent: @patch("litellm.proxy.proxy_server.prisma_client") def test_append_event_no_status_update_for_unknown_type(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -357,6 +368,7 @@ class TestAppendWorkflowEvent: @patch("litellm.proxy.proxy_server.prisma_client") def test_sequence_number_increments(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -377,6 +389,7 @@ class TestAppendWorkflowEvent: @patch("litellm.proxy.proxy_server.prisma_client") def test_unknown_run_id_returns_404(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None) resp = self.client.post( @@ -389,6 +402,7 @@ class TestAppendWorkflowEvent: def test_sequence_collision_retries_and_succeeds(self, mock_pc): """UniqueViolationError on first attempt triggers retry; second attempt succeeds.""" mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -427,6 +441,7 @@ class TestWorkflowMessages: @patch("litellm.proxy.proxy_server.prisma_client") def test_append_message(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -444,6 +459,7 @@ class TestWorkflowMessages: @patch("litellm.proxy.proxy_server.prisma_client") def test_append_message_unknown_run_returns_404(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None) resp = self.client.post( @@ -455,6 +471,7 @@ class TestWorkflowMessages: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_messages_ordered(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -475,6 +492,7 @@ class TestWorkflowMessages: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_messages_respects_limit(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -498,6 +516,7 @@ class TestListWorkflowEvents: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_events_ordered(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -518,6 +537,7 @@ class TestListWorkflowEvents: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_events_respects_limit(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run() ) @@ -531,6 +551,7 @@ class TestListWorkflowEvents: @patch("litellm.proxy.proxy_server.prisma_client") def test_list_events_unknown_run_returns_404(self, mock_pc): mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None) resp = self.client.get("/v1/workflows/runs/nonexistent/events") @@ -553,6 +574,7 @@ class TestTenantIsolation: token = "tok-owner" client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token)) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.create = AsyncMock( return_value=_make_run(created_by=token) ) @@ -567,6 +589,7 @@ class TestTenantIsolation: token = "tok-owner" client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token)) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[]) resp = client.get("/v1/workflows/runs") @@ -578,6 +601,7 @@ class TestTenantIsolation: def test_admin_list_not_scoped(self, mock_pc): client = self._make_app_with_auth(_override_auth_admin) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[]) resp = client.get("/v1/workflows/runs") @@ -590,6 +614,7 @@ class TestTenantIsolation: token = "tok-caller" client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token)) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db # Run owned by a different key self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by="tok-other-owner") @@ -603,6 +628,7 @@ class TestTenantIsolation: token = "tok-caller" client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token)) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by=None) ) @@ -615,6 +641,7 @@ class TestTenantIsolation: token = "tok-caller" client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token)) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by=None) ) @@ -631,6 +658,7 @@ class TestTenantIsolation: token = "tok-caller" client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token)) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by=token) ) @@ -660,6 +688,7 @@ class TestAdminViewerReadParity: def test_admin_viewer_list_not_scoped(self, mock_pc): client = self._make_app_with_auth(_override_auth_admin_viewer) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[]) resp = client.get("/v1/workflows/runs") @@ -671,6 +700,7 @@ class TestAdminViewerReadParity: def test_admin_viewer_get_other_owners_run_succeeds(self, mock_pc): client = self._make_app_with_auth(_override_auth_admin_viewer) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by="tok-other-owner") ) @@ -682,6 +712,7 @@ class TestAdminViewerReadParity: def test_admin_viewer_lists_other_owners_events(self, mock_pc): client = self._make_app_with_auth(_override_auth_admin_viewer) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by="tok-other-owner") ) @@ -697,6 +728,7 @@ class TestAdminViewerReadParity: def test_admin_viewer_lists_other_owners_messages(self, mock_pc): client = self._make_app_with_auth(_override_auth_admin_viewer) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by="tok-other-owner") ) @@ -713,6 +745,7 @@ class TestAdminViewerReadParity: """Read parity must not become write parity: PATCH still passes the caller through.""" client = self._make_app_with_auth(_override_auth_admin_viewer) mock_pc.db = self._prisma.db + mock_pc.replica_db = mock_pc.db self._prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by="tok-other-owner") ) @@ -730,6 +763,7 @@ class TestAdminViewerReadParity: prisma.db.litellm_workflowrun.find_unique = AsyncMock( return_value=_make_run(created_by="tok-other-owner") ) + prisma.replica_db = prisma.db with pytest.raises(HTTPException) as exc_info: asyncio.run(_require_run(prisma, "run-1", _override_auth_admin_viewer())) diff --git a/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py b/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py index b1d111bf1f9..81c26e71310 100644 --- a/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py +++ b/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py @@ -192,6 +192,7 @@ class TestCreateAuditLogForUpdateWithCallbacks: patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, ): mock_prisma.db.litellm_auditlog.create = AsyncMock() + mock_prisma.replica_db = mock_prisma.db audit_log = _make_audit_log() await create_audit_log_for_update(audit_log) @@ -219,6 +220,7 @@ class TestCreateAuditLogForUpdateWithCallbacks: mock_logger.async_log_audit_log_event.assert_not_called() mock_prisma.db.litellm_auditlog.create.assert_not_called() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.asyncio async def test_no_dispatch_when_store_audit_logs_false(self, monkeypatch: pytest.MonkeyPatch): @@ -267,6 +269,7 @@ class TestCreateAuditLogForUpdateWithCallbacks: mock_prisma.db.litellm_auditlog.create = AsyncMock( side_effect=RuntimeError("DB connection lost") ) + mock_prisma.replica_db = mock_prisma.db audit_log = _make_audit_log() await create_audit_log_for_update(audit_log) diff --git a/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py b/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py index b5349fc2387..40478e6c6b5 100644 --- a/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py +++ b/tests/test_litellm/proxy/management_helpers/test_bulk_user_creation.py @@ -147,6 +147,7 @@ class _FakePrisma: raced_ids: frozenset[str] = frozenset(), ) -> None: self.db = _Db(teams or [], fail_ids, commit_then_drop, raced_ids) + self.replica_db = self.db self.tx_count = 0 self.locks: list[str] = [] @@ -252,6 +253,7 @@ async def test_one_insert_and_one_locked_write_per_team(): async def test_bad_rows_fail_alone_and_good_rows_still_land(): prisma = _FakePrisma(teams=[_team("t1")]) prisma.db.litellm_usertable.rows["taken"] = _UserRow(user_id="taken", user_email="Taken@Example.com") + prisma.replica_db = prisma.db response = await _run( prisma, [ @@ -325,6 +327,7 @@ async def test_team_write_failure_keeps_user_and_reports_it_on_the_row(): raise RuntimeError("roster write failed") prisma.db.litellm_teamtable.update = explode + prisma.replica_db = prisma.db response = await _run(prisma, [{"user_id": "u1", "teams": ["t1", "t2"]}]) result = response.data[0] @@ -394,6 +397,7 @@ async def test_non_admin_cannot_create_admin_users_but_other_rows_proceed(): async def test_license_is_checked_once_against_the_whole_batch(): prisma = _FakePrisma() prisma.db.litellm_usertable.rows["existing"] = _UserRow(user_id="existing") + prisma.replica_db = prisma.db license = _License(max_users=3) with pytest.raises(ManagementProblem) as exc: diff --git a/tests/test_litellm/proxy/management_helpers/test_bulk_user_deletion.py b/tests/test_litellm/proxy/management_helpers/test_bulk_user_deletion.py index 32e5bea613c..a3fb8f85db3 100644 --- a/tests/test_litellm/proxy/management_helpers/test_bulk_user_deletion.py +++ b/tests/test_litellm/proxy/management_helpers/test_bulk_user_deletion.py @@ -173,6 +173,7 @@ class _FakePrisma: fail_commit: bool = False, ) -> None: self.db = _Db(users, teams, memberships, tokens, invitations, org_memberships, jwt_mappings) + self.replica_db = self.db self._on_lock = on_lock self._fail_locks = fail_locks self._fail_commit = fail_commit @@ -189,6 +190,7 @@ class _FakePrisma: raise RuntimeError("connection reset") except BaseException: self.db.__dict__.update(snapshot.__dict__) + self.replica_db = self.db raise self.locks.extend(tx.locks) self.roster_reads.extend(tx.roster_reads) diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index 922504ecc58..cd5a6a991a4 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -193,6 +193,7 @@ async def test_add_new_member_links_default_team_budget_id(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -265,6 +266,7 @@ async def test_add_new_member_no_budget_when_default_budget_row_is_missing(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -318,6 +320,7 @@ async def test_add_new_member_budget_duration_only_clones_default_max_budget(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -403,6 +406,7 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -494,6 +498,7 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -572,6 +577,7 @@ async def test_add_new_member_persists_budget_duration(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -636,6 +642,7 @@ async def test_add_new_member_persists_budget_duration_without_max_budget(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -709,6 +716,7 @@ async def test_add_new_member_with_user_email_links_default_budget(): mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock( return_value=mock_default_budget_row ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.create = AsyncMock() mock_team_membership_response = MagicMock() @@ -845,6 +853,7 @@ async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): db: Final = _FakeDb() prisma_client: Final = MagicMock() prisma_client.db = db + prisma_client.replica_db = prisma_client.db admin: Final = UserAPIKeyAuth(user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN) team_id: Final = "team-shared-default" default_budget: Final = await db.litellm_budgettable.create(data={"budget_id": "team-default", "max_budget": 100.0}) @@ -948,6 +957,7 @@ async def test_attach_object_permission_to_dict_with_object_permission_id(): mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( return_value=mock_object_permission ) + mock_prisma_client.replica_db = mock_prisma_client.db # Call the function result = await attach_object_permission_to_dict( @@ -994,6 +1004,7 @@ async def test_attach_object_permission_to_dict_without_object_permission_id(): # Verify no database query was made mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_not_called() + mock_prisma_client.replica_db = mock_prisma_client.db @pytest.mark.asyncio @@ -1021,6 +1032,7 @@ async def test_attach_object_permission_to_dict_object_permission_not_found(): mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( return_value=None ) + mock_prisma_client.replica_db = mock_prisma_client.db # Call the function result = await attach_object_permission_to_dict( @@ -1075,6 +1087,7 @@ async def test_attach_object_permission_to_dict_with_dict_method(): mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( return_value=mock_object_permission ) + mock_prisma_client.replica_db = mock_prisma_client.db # Call the function result = await attach_object_permission_to_dict( @@ -1138,6 +1151,7 @@ async def test_attach_object_permission_to_dict_with_empty_dict(): # Verify no database query was made mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_not_called() + mock_prisma_client.replica_db = mock_prisma_client.db @pytest.mark.asyncio @@ -1170,6 +1184,7 @@ async def test_attach_object_permission_to_dict_with_none_object_permission_id() # Verify no database query was made mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_not_called() + mock_prisma_client.replica_db = mock_prisma_client.db @pytest.mark.asyncio @@ -1202,6 +1217,7 @@ async def test_add_new_member_appends_team_only_if_absent_for_existing_user(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_after) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) mock_membership = MagicMock() @@ -1275,6 +1291,7 @@ async def test_add_new_member_creates_missing_user_atomically_via_upsert(): "user_role": "internal_user", } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_created) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() mock_prisma_client.db.litellm_usertable.create = AsyncMock() mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) @@ -1383,5 +1400,6 @@ async def test_add_new_member_runs_every_write_on_the_caller_transaction(new_mem assert tx.litellm_usertable.upsert.await_count + tx.litellm_usertable.create.await_count == 1 prisma_client.db.assert_not_called() + prisma_client.replica_db = prisma_client.db prisma_client.get_data.assert_not_awaited() prisma_client.insert_data.assert_not_awaited() diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index 078315c2bf8..c1573dbc69e 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -42,6 +42,7 @@ async def test_set_object_permission(): mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=mock_created_permission ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) # Test data with object_permission @@ -104,6 +105,7 @@ async def test_set_object_permission_persists_mcp_tool_search_enabled(): mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=mock_created_permission ) + mock_prisma_client.replica_db = mock_prisma_client.db data_json = { "object_permission": { @@ -130,6 +132,7 @@ async def test_set_object_permission_persists_skills(): mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=mock_created_permission ) + mock_prisma_client.replica_db = mock_prisma_client.db data_json = { "object_permission": LiteLLM_ObjectPermissionBase(skills=["private-skill"]).model_dump(), @@ -778,6 +781,7 @@ async def test_validate_db_mcp_server_alias_outside_team_scope_raises_when_regis mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock( return_value=[mock_db_server] ) + mock_prisma_client.replica_db = mock_prisma_client.db team_obj = _make_team_obj(mcp_servers=[]) with pytest.raises(HTTPException) as exc_info: @@ -1266,6 +1270,7 @@ def _make_grandfather_fixtures(mcp_servers=None, mcp_tool_permissions=None): existing_row.mcp_tool_permissions = mcp_tool_permissions or {} mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db return mock_prisma, existing_row @@ -1392,6 +1397,7 @@ def _make_ambiguity_prisma(existing_tool_permissions=None): object permission row (if any) stores the given mcp_tool_permissions JSON string.""" mock_prisma = MagicMock() mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(_SHARED_ALIAS_DB_SERVERS)) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_objectpermissiontable.create = AsyncMock( return_value=MagicMock(object_permission_id="perm-id") ) @@ -1423,6 +1429,7 @@ async def test_set_object_permission_rejects_shared_alias_or_name_tool_permissio assert exc_info.value.status_code == 400 assert all(server_id in str(exc_info.value.detail) for server_id in colliding_ids) mock_prisma.db.litellm_objectpermissiontable.create.assert_not_called() + mock_prisma.replica_db = mock_prisma.db @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/management_helpers/test_team_metadata_validation.py b/tests/test_litellm/proxy/management_helpers/test_team_metadata_validation.py index dfb834dc31f..ce48b384dd8 100644 --- a/tests/test_litellm/proxy/management_helpers/test_team_metadata_validation.py +++ b/tests/test_litellm/proxy/management_helpers/test_team_metadata_validation.py @@ -364,6 +364,7 @@ async def _drive_create(metadata, mock_sink=None): pc.get_data = AsyncMock(return_value=None) pc.update_data = AsyncMock(return_value=MagicMock()) pc.db.litellm_teamtable.create = AsyncMock(return_value=team_row) + pc.replica_db = pc.db pc.db.litellm_teamtable.count = AsyncMock(return_value=0) pc.db.litellm_teamtable.update = AsyncMock(return_value=team_row) pc.db.litellm_usertable.update = AsyncMock(return_value=MagicMock()) @@ -410,6 +411,7 @@ async def _drive_update(kind, existing_metadata, payload): ), ): pc.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing) + pc.replica_db = pc.db pc.db.litellm_teamtable.update = AsyncMock( return_value=LiteLLM_TeamTable(team_id=team_id, team_alias="matrix") ) diff --git a/tests/test_litellm/proxy/memory/test_memory_endpoints.py b/tests/test_litellm/proxy/memory/test_memory_endpoints.py index dff0e80fa77..7e805a94b9e 100644 --- a/tests/test_litellm/proxy/memory/test_memory_endpoints.py +++ b/tests/test_litellm/proxy/memory/test_memory_endpoints.py @@ -212,6 +212,7 @@ def _make_team(team_id: str, *, admin_user_ids: List[str]) -> Any: def _make_prisma() -> MagicMock: client = MagicMock() client.db = MagicMock() + client.replica_db = client.db client.db.litellm_memorytable = _InMemoryMemoryTable() client.db.litellm_teamtable = _InMemoryTeamTable() return client diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index 50768e48d43..f381ef0f7ba 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -54,6 +54,7 @@ async def test_map_raw_file_ids_to_unified_empty_ids_skips_db(): assert await map_raw_file_ids_to_unified(frozenset(), prisma_client) == {} prisma_client.db.litellm_managedfiletable.find_many.assert_not_called() + prisma_client.replica_db = prisma_client.db @pytest.mark.asyncio @@ -70,6 +71,7 @@ async def test_map_raw_file_ids_to_unified_bulk_queries_and_filters_to_requested row_b = MagicMock(unified_file_id="unified-b", flat_model_file_ids=["file-raw-b"]) prisma_client = MagicMock() prisma_client.db.litellm_managedfiletable.find_many = AsyncMock(return_value=[row_a, row_b]) + prisma_client.replica_db = prisma_client.db mapping = await map_raw_file_ids_to_unified( frozenset({"file-raw-b", "file-raw-a", "file-raw-missing"}), prisma_client @@ -204,6 +206,7 @@ async def _run_update(monkeypatch, poller_active: bool) -> dict: prisma_client = MagicMock() update_mock = AsyncMock() prisma_client.db.litellm_managedobjecttable.update = update_mock + prisma_client.replica_db = prisma_client.db db_batch_object = MagicMock() db_batch_object.status = "in_progress" @@ -285,6 +288,7 @@ async def test_retrieving_a_batch_whose_status_is_unchanged_writes_nothing(monke prisma_client = MagicMock() update_mock = AsyncMock() prisma_client.db.litellm_managedobjecttable.update = update_mock + prisma_client.replica_db = prisma_client.db db_batch_object = MagicMock() db_batch_object.status = "completed" @@ -310,6 +314,7 @@ async def test_update_batch_in_database_is_a_noop_for_unmanaged_batches(monkeypa prisma_client = MagicMock() update_mock = AsyncMock() prisma_client.db.litellm_managedobjecttable.update = update_mock + prisma_client.replica_db = prisma_client.db await cu.update_batch_in_database( batch_id="batch-raw-xyz", @@ -334,6 +339,7 @@ async def test_the_caller_s_accounting_decision_wins_over_a_later_poller_transit prisma_client = MagicMock() update_mock = AsyncMock() prisma_client.db.litellm_managedobjecttable.update = update_mock + prisma_client.replica_db = prisma_client.db db_batch_object = MagicMock() db_batch_object.status = "in_progress" @@ -364,6 +370,7 @@ async def test_a_caller_that_handed_off_accounting_still_leaves_the_marker_alone prisma_client = MagicMock() update_mock = AsyncMock() prisma_client.db.litellm_managedobjecttable.update = update_mock + prisma_client.replica_db = prisma_client.db db_batch_object = MagicMock() db_batch_object.status = "in_progress" diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 48699b47e7f..5612a76a6fd 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -3984,6 +3984,7 @@ def test_get_file_content_model_routed_attaches_trusted_model_credentials(monkey managed_file_row.storage_url = None prisma_stub = MagicMock() prisma_stub.db.litellm_managedfiletable.find_first = AsyncMock(return_value=managed_file_row) + prisma_stub.replica_db = prisma_stub.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_stub) setup_proxy_logging_object(monkeypatch, router) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py b/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py index dc8c49b93d8..817d69d95d1 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_managed_id_rewriter.py @@ -20,6 +20,7 @@ def _user() -> UserAPIKeyAuth: def _prisma_client(file_rows=None, batch_rows=None) -> MagicMock: pc = MagicMock() pc.db = MagicMock() + pc.replica_db = pc.db pc.db.litellm_managedfiletable = MagicMock() pc.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) pc.db.litellm_managedfiletable.find_many = AsyncMock( @@ -118,6 +119,7 @@ async def test_list_batches_out_of_range_limit_raises_400(limit, expected_messag assert exc.value.openai_code == expected_openai_code assert exc.value.message == expected_message pc.db.litellm_managedobjecttable.find_many.assert_not_called() + pc.replica_db = pc.db @pytest.mark.asyncio @@ -140,6 +142,7 @@ async def test_list_batches_limit_zero_returns_empty_page_without_db_query(): "has_more": False, } pc.db.litellm_managedobjecttable.find_many.assert_not_called() + pc.replica_db = pc.db @pytest.mark.asyncio @@ -199,6 +202,7 @@ async def test_streamed_response_is_owned_and_rewritten_across_chunk_boundaries( ) pc.db.litellm_managedobjecttable.upsert.assert_awaited_once() + pc.replica_db = pc.db created = pc.db.litellm_managedobjecttable.upsert.await_args.kwargs["data"]["create"] assert created["created_by"] == "user-1" assert created["team_id"] == "team-1" @@ -229,6 +233,7 @@ async def test_streamed_response_with_cr_only_frame_delimiters_is_still_owned_an ) pc.db.litellm_managedobjecttable.upsert.assert_awaited_once() + pc.replica_db = pc.db managed_id = pc.db.litellm_managedobjecttable.upsert.await_args.kwargs["data"]["create"]["unified_object_id"] assert RAW_RESPONSE_ID.encode() not in output assert output == _response_stream_bytes(managed_id).replace(b"\n", b"\r") @@ -252,12 +257,14 @@ async def test_streamed_bytes_untouched_on_routes_without_a_response_id(): assert output == payload pc.db.litellm_managedobjecttable.upsert.assert_not_awaited() + pc.replica_db = pc.db @pytest.mark.asyncio async def test_streamed_response_stays_raw_and_intact_when_the_row_cannot_be_persisted(): pc = _prisma_client() pc.db.litellm_managedobjecttable.upsert = AsyncMock(side_effect=RuntimeError("db down")) + pc.replica_db = pc.db payload = _response_stream_bytes() output = await _collect( diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py b/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py index 1ca830dc1e6..f18729d5150 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_engine_endpoints.py @@ -103,6 +103,7 @@ class TestListPoliciesIncludesConfig: row = _make_policy_row(policy_id="uuid-1", policy_name="db-policy", guardrails_add=["db-guard"]) prisma = MagicMock() prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + prisma.replica_db = prisma.db _set_prisma(monkeypatch, prisma) policy_registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) @@ -124,6 +125,7 @@ class TestListPoliciesIncludesConfig: row = _make_policy_row(policy_id="uuid-1", policy_name="shared-name", guardrails_add=["db-guard"]) prisma = MagicMock() prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + prisma.replica_db = prisma.db _set_prisma(monkeypatch, prisma) policy_registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) @@ -146,6 +148,7 @@ class TestListPoliciesIncludesConfig: ) prisma = MagicMock() prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + prisma.replica_db = prisma.db _set_prisma(monkeypatch, prisma) policy_registry.load_policies({"shared-name": {"guardrails": {"add": ["config-guard"]}}}) @@ -171,11 +174,13 @@ class TestListPoliciesIncludesConfig: production_row = _make_policy_row(policy_id="uuid-1", policy_name="shared-name", guardrails_add=["db-guard"]) sync_prisma = MagicMock() sync_prisma.db.litellm_policytable.find_many = AsyncMock(side_effect=[[production_row], []]) + sync_prisma.replica_db = sync_prisma.db await policy_registry.sync_policies_from_db(sync_prisma) assert policy_registry.get_source("shared-name") == "db" fresh_prisma = MagicMock() fresh_prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[]) + fresh_prisma.replica_db = fresh_prisma.db _set_prisma(monkeypatch, fresh_prisma) response = await policy_endpoints.list_policies() @@ -191,6 +196,7 @@ class TestListPoliciesIncludesConfig: row = _make_policy_row(policy_id="uuid-1", policy_name="db-policy", version_status="draft") prisma = MagicMock() prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[row]) + prisma.replica_db = prisma.db _set_prisma(monkeypatch, prisma) policy_registry.load_policies({"config-policy": {"guardrails": {"add": ["tooling"]}}}) @@ -232,6 +238,7 @@ class TestListAttachmentsIncludesConfig: row = _make_attachment_row(attachment_id="att-1", policy_name="db-policy") prisma = MagicMock() prisma.db.litellm_policyattachmenttable.find_many = AsyncMock(return_value=[row]) + prisma.replica_db = prisma.db _set_prisma(monkeypatch, prisma) attachment_registry.load_attachments([{"policy": "config-policy", "scope": "*"}]) diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_validator.py b/tests/test_litellm/proxy/policy_engine/test_policy_validator.py index a56695ee7f5..075736717ca 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_validator.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_validator.py @@ -41,6 +41,7 @@ class _FakePrisma: def __init__(self, teams: Set[str] = frozenset(), keys: Set[str] = frozenset()): self.db = _FakeDB(teams, keys) + self.replica_db = self.db class _FakeRouter: diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py index b6633779326..d37a0d988f8 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning.py @@ -112,6 +112,7 @@ class TestSyncPoliciesFromDbProductionOnly: prisma = MagicMock() prod_row = _make_row(policy_id="prod-1", version_status="production") prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + prisma.replica_db = prisma.db result = await registry.get_all_policies_from_db( prisma, version_status="production" @@ -134,6 +135,7 @@ class TestSyncPoliciesFromDbProductionOnly: guardrails_add=["g1"], ) prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + prisma.replica_db = prisma.db await registry.sync_policies_from_db(prisma) @@ -156,6 +158,7 @@ class TestUpdatePolicyDraftOnly: prisma = MagicMock() prod_row = _make_row(policy_id="pid-1", version_status="production") prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=prod_row) + prisma.replica_db = prisma.db with pytest.raises(Exception, match='Error updating policy in DB: Only draft versions can be') as exc_info: await registry.update_policy_in_db( @@ -187,6 +190,7 @@ class TestUpdatePolicyDraftOnly: description="new", ) prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft_row) + prisma.replica_db = prisma.db prisma.db.litellm_policytable.update = AsyncMock(return_value=updated_row) result = await registry.update_policy_in_db( @@ -215,6 +219,7 @@ class TestDeletePolicyFromDb: version_status="production", ) prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=prod_row) + prisma.replica_db = prisma.db prisma.db.litellm_policytable.delete = AsyncMock() result = await registry.delete_policy_from_db( @@ -238,6 +243,7 @@ class TestDeletePolicyFromDb: version_status="draft", ) prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft_row) + prisma.replica_db = prisma.db prisma.db.litellm_policytable.delete = AsyncMock() result = await registry.delete_policy_from_db( @@ -268,6 +274,7 @@ class TestCreateNewVersion: ) # find_first for production prisma.db.litellm_policytable.find_first = AsyncMock(return_value=prod) + prisma.replica_db = prisma.db # find_first for latest version number prisma.db.litellm_policytable.find_first.side_effect = [ prod, # production lookup @@ -321,6 +328,7 @@ class TestUpdateVersionStatus: published_at=datetime.now(timezone.utc), ) prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft) + prisma.replica_db = prisma.db prisma.db.litellm_policytable.update = AsyncMock(return_value=updated) result = await registry.update_version_status( @@ -340,6 +348,7 @@ class TestUpdateVersionStatus: prisma = MagicMock() draft = _make_row(policy_id="d-1", version_status="draft") prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=draft) + prisma.replica_db = prisma.db with pytest.raises(Exception, match='Error updating version status: Cannot promote draft') as exc_info: await registry.update_version_status( @@ -370,6 +379,7 @@ class TestUpdateVersionStatus: prisma.db.litellm_policytable.find_unique = AsyncMock( return_value=published_row ) + prisma.replica_db = prisma.db prisma.db.litellm_policytable.update_many = AsyncMock() prisma.db.litellm_policytable.update = AsyncMock(return_value=updated_row) @@ -406,6 +416,7 @@ class TestCompareVersions: guardrails_add=["g1", "g2"], ) prisma.db.litellm_policytable.find_unique = AsyncMock(side_effect=[a, b]) + prisma.replica_db = prisma.db result = await registry.compare_versions( policy_id_a="a", @@ -434,6 +445,7 @@ class TestResolveGuardrailsProductionOnly: guardrails_add=["g1"], ) prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + prisma.replica_db = prisma.db result = await registry.resolve_guardrails_from_db( policy_name="base", @@ -457,6 +469,7 @@ class TestGetPolicyRegistrySingleton: def _prisma_with_policy_rows(production_rows, non_production_rows=None): prisma = MagicMock() prisma.db.litellm_policytable.find_many = AsyncMock(side_effect=[production_rows, non_production_rows or []]) + prisma.replica_db = prisma.db return prisma @@ -595,6 +608,7 @@ class TestRemovePolicyRestoresConfigFallback: prisma = MagicMock() prod_row = _make_row(policy_id="prod-1", policy_name="shared-name", version_status="production") prisma.db.litellm_policytable.find_unique = AsyncMock(return_value=prod_row) + prisma.replica_db = prisma.db prisma.db.litellm_policytable.delete = AsyncMock() result = await registry.delete_policy_from_db(policy_id="prod-1", prisma_client=prisma) @@ -612,6 +626,7 @@ class TestRemovePolicyRestoresConfigFallback: registry.add_policy("shared-name", Policy(guardrails=PolicyGuardrails(add=["db-guard"])), source="db") prisma = MagicMock() prisma.db.litellm_policytable.delete_many = AsyncMock() + prisma.replica_db = prisma.db result = await registry.delete_all_versions(policy_name="shared-name", prisma_client=prisma) @@ -626,6 +641,7 @@ class TestRemovePolicyRestoresConfigFallback: registry.add_policy("db-only", Policy(guardrails=PolicyGuardrails(add=["db-guard"])), source="db") prisma = MagicMock() prisma.db.litellm_policytable.delete_many = AsyncMock() + prisma.replica_db = prisma.db result = await registry.delete_all_versions(policy_name="db-only", prisma_client=prisma) diff --git a/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py b/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py index 9764bc2e465..21972cec245 100644 --- a/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py +++ b/tests/test_litellm/proxy/policy_engine/test_policy_versioning_e2e.py @@ -77,6 +77,7 @@ async def test_full_lifecycle_create_draft_edit_publish_promote(): return created_v1 prisma.db.litellm_policytable.create = AsyncMock(side_effect=create_impl) + prisma.replica_db = prisma.db req = PolicyCreateRequest( policy_name="lifecycle-policy", description="Initial", @@ -205,6 +206,7 @@ async def test_attachments_resolve_against_production_after_promotion(): guardrails_add=["ga", "gb"], ) prisma.db.litellm_policytable.find_many = AsyncMock(return_value=[prod_row]) + prisma.replica_db = prisma.db resolved = await registry.resolve_guardrails_from_db( policy_name="att-policy", diff --git a/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py b/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py index 41c3003af34..06d6b02794f 100644 --- a/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py +++ b/tests/test_litellm/proxy/prompts/test_prompt_endpoints.py @@ -396,6 +396,7 @@ class TestConfigPromptInfoWithEnvironment: mock_prisma = MagicMock() mock_prisma.db.litellm_prompttable.find_many = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db return mock_prisma @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py b/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py index 387448bb9b3..f483e297b26 100644 --- a/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py +++ b/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py @@ -47,6 +47,7 @@ async def test_delete_prompt_success(): # Mock DB Client mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_prompttable.delete_many = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock In-Memory Registry with patch( @@ -102,6 +103,7 @@ async def test_delete_prompt_by_base_id_success(): # Mock DB Client mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_prompttable.delete_many = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db # Mock In-Memory Registry with patch( @@ -147,6 +149,7 @@ async def test_delete_prompt_environment_scope_reaches_db_and_registry(): mock_user_auth = UserAPIKeyAuth(api_key="sk-1234", user_role=LitellmUserRoles.PROXY_ADMIN) mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_prompttable.delete_many = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch( # test-quality-ok: stubs the collaborator so the test pins what the endpoint deletes "litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY" @@ -235,6 +238,7 @@ async def test_patch_prompt_row_deleted_mid_update_returns_404(): mock_prisma_client.db.litellm_prompttable.find_many = AsyncMock( return_value=[target_row] ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_prompttable.update = AsyncMock(return_value=None) existing_prompt = PromptSpec( @@ -275,6 +279,7 @@ async def test_patch_prompt_merges_unsent_fields_from_db_row_not_stale_memory(): db_row = _db_row("Begin every reply with HOWDY") mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[db_row]) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_prompttable.update = AsyncMock(return_value=db_row) stale_in_memory = PromptSpec( prompt_id="test_prompt.v1", @@ -440,6 +445,7 @@ async def test_patch_prompt_info_only_keeps_legacy_keyed_row_patchable(): mock_prisma_client.db.litellm_prompttable.find_many = AsyncMock( return_value=[target_row] ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_prompttable.update = AsyncMock(return_value=updated_row) existing_prompt = PromptSpec( diff --git a/tests/test_litellm/proxy/prompts/test_prompt_environment.py b/tests/test_litellm/proxy/prompts/test_prompt_environment.py index 3cb647dea13..66a59c9e95b 100644 --- a/tests/test_litellm/proxy/prompts/test_prompt_environment.py +++ b/tests/test_litellm/proxy/prompts/test_prompt_environment.py @@ -109,6 +109,7 @@ async def test_create_prompt_stores_environment_and_created_by(): mock_prisma_client.db.litellm_prompttable.create = AsyncMock( return_value=mock_db_entry ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[]) request = Prompt( @@ -157,6 +158,7 @@ async def test_update_prompt_stores_environment_and_created_by(): mock_prisma_client.db.litellm_prompttable.find_many = AsyncMock( return_value=[mock_existing] ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_db_entry = MagicMock() mock_db_entry.model_dump.return_value = { @@ -223,6 +225,7 @@ async def test_delete_prompt_scoped_to_environment(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_prompttable.delete_many = AsyncMock(return_value=None) + mock_prisma_client.replica_db = mock_prisma_client.db with patch( "litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY" diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 36d2e16d261..ecd00f692a9 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -925,6 +925,7 @@ async def test_tuning_baseline_v3_is_created_alongside_the_legacy_row(): prisma_client = MagicMock() prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_config.create = AsyncMock() deployment = { "model_name": "a", @@ -962,6 +963,7 @@ async def test_scorer_baseline_upgrade_preserves_existing_routers_and_is_not_ref else None ) ) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_config.create = AsyncMock() baseline = await ProxyStartupEvent._load_heuristic_v1_tuning_baselines(prisma_client, deployments) @@ -1003,6 +1005,7 @@ async def test_tuning_baseline_waits_for_a_complete_db_model_census(monkeypatch) assert result is None prisma_client.db.litellm_config.find_unique.assert_not_called() + prisma_client.replica_db = prisma_client.db # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py index 40bd66ea91c..00a1464fc58 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_cost_map.py @@ -45,6 +45,7 @@ def _attach_litellm_config(mock_prisma): table.delete = AsyncMock() table.delete_many = AsyncMock() mock_prisma.db.litellm_config = table + mock_prisma.replica_db = mock_prisma.db return table diff --git a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py index 0731c233fef..3d0b750823d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_spend_counters.py +++ b/tests/test_litellm/proxy/proxy_server/test_spend_counters.py @@ -239,6 +239,7 @@ def _make_prisma_with_end_user_row(spend: float | None): prisma.db.litellm_endusertable.find_unique = AsyncMock( return_value=None if spend is None else MagicMock(spend=spend) ) + prisma.replica_db = prisma.db return prisma @@ -262,6 +263,7 @@ async def test_get_current_spend_end_user_floor_admits_after_a_reset_on_a_stale_ assert result == 0.0 prisma.db.litellm_endusertable.find_unique.assert_awaited_once_with(where={"user_id": "customer-42"}) + prisma.replica_db = prisma.db fake_cache.redis_cache.async_set_max.assert_not_called() @@ -331,6 +333,7 @@ async def test_get_current_spend_floors_window_against_spend_logs(monkeypatch): def _make_window_spend_prisma(row=None, spend_logs_total=0.0): prisma = MagicMock() prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=row) + prisma.replica_db = prisma.db prisma.db.litellm_spendlogs.group_by = AsyncMock( return_value=[{"api_key": "tok", "_sum": {"spend": spend_logs_total}}] ) @@ -366,6 +369,7 @@ async def test_get_current_spend_floors_window_against_maintained_row(monkeypatc assert result == 15.0 fake_prisma.db.litellm_spendlogs.group_by.assert_not_awaited() + fake_prisma.replica_db = fake_prisma.db fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key=counter_key, value=15.0) @@ -397,6 +401,7 @@ async def test_get_current_spend_floors_window_against_logs_when_row_stale(monke assert result == 15.0 fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once() + fake_prisma.replica_db = fake_prisma.db @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py index 2d654ea28ec..ff6a385160c 100644 --- a/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py +++ b/tests/test_litellm/proxy/rag_endpoints/test_rag_endpoints.py @@ -552,6 +552,7 @@ def test_rag_ingest_rejects_non_string_provider(client_internal_user): def test_rag_ingest_never_creates_db_row_for_registry_store(client_internal_user): prisma_client = MagicMock() prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db create_in_db = AsyncMock() aingest_patch, registry_patch = _patched_ingest_boundary( S3_REGISTRY_STORE, {"vector_store_id": "s3-store", "file_id": "file_123"} @@ -576,6 +577,7 @@ def test_rag_ingest_never_creates_db_row_for_registry_store(client_internal_user def test_rag_ingest_fresh_store_creates_db_row_with_the_requesters_params(client_internal_user): prisma_client = MagicMock() prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db create_in_db = AsyncMock() with ( patch( # test-quality-ok: aingest is the endpoint's downstream boundary; persistence is what the test asserts @@ -633,6 +635,7 @@ async def test_save_vector_store_from_rag_ingest_appends_file_to_db_managed_stor existing_row.vector_store_metadata = {"ingested_files": [{"file_id": "file_old"}]} prisma_client = MagicMock() table = prisma_client.db.litellm_managedvectorstorestable + prisma_client.replica_db = prisma_client.db table.find_unique = AsyncMock(return_value=existing_row) table.update = AsyncMock() create_in_db = AsyncMock() @@ -660,6 +663,7 @@ async def test_save_vector_store_from_rag_ingest_still_creates_row_for_fresh_sto prisma_client = MagicMock() prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db create_in_db = AsyncMock() with patch( # test-quality-ok: the DB write boundary whose inputs the test asserts diff --git a/tests/test_litellm/proxy/spend_tracking/test_cloudzero_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_cloudzero_endpoints.py index 45aa065380d..e6e1e8c0fc3 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_cloudzero_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_cloudzero_endpoints.py @@ -30,6 +30,7 @@ async def test_delete_cloudzero_settings_success(client, monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config = mock_litellm_config monkeypatch.setattr(ps, "prisma_client", mock_prisma) @@ -57,6 +58,7 @@ async def test_delete_cloudzero_settings_not_found(client, monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config = mock_litellm_config monkeypatch.setattr(ps, "prisma_client", mock_prisma) @@ -93,6 +95,7 @@ async def test_get_cloudzero_settings_success(client, monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config = mock_litellm_config monkeypatch.setattr(ps, "prisma_client", mock_prisma) @@ -134,6 +137,7 @@ async def test_get_cloudzero_settings_not_configured(client, monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config = mock_litellm_config monkeypatch.setattr(ps, "prisma_client", mock_prisma) @@ -168,6 +172,7 @@ async def test_get_cloudzero_settings_empty_param_value(client, monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config = mock_litellm_config monkeypatch.setattr(ps, "prisma_client", mock_prisma) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py index 3e4b817fab8..6c9e0a0a658 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_counter_batch.py @@ -288,6 +288,7 @@ async def test_batch_reads_never_touch_a_prisma_client_when_redis_answers(monkey redis = CountingRedis({"spend:key:hashed": 3.0}) prisma = MagicMock() prisma.db.litellm_verificationtoken.find_unique = AsyncMock() + prisma.replica_db = prisma.db monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) monkeypatch.setattr(ps, "prisma_client", prisma) @@ -301,6 +302,7 @@ async def test_batch_reads_never_touch_a_prisma_client_when_redis_answers(monkey def _reseed_prisma(spend: float) -> MagicMock: prisma = MagicMock() prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=MagicMock(spend=spend)) + prisma.replica_db = prisma.db return prisma @@ -381,6 +383,7 @@ async def test_post_call_cold_counters_seed_from_the_mget_miss_without_a_second_ redis.async_set_cache = AsyncMock(return_value=True) prisma = MagicMock() prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=MagicMock(spend=4.0)) + prisma.replica_db = prisma.db monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) monkeypatch.setattr(ps, "prisma_client", prisma) @@ -444,6 +447,7 @@ async def test_reconcile_settles_a_flushed_counter_on_its_own_after_the_shared_p redis.async_set_max = AsyncMock(return_value=4.0) prisma = MagicMock() prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=MagicMock(spend=4.0)) + prisma.replica_db = prisma.db monkeypatch.setattr(ps, "spend_counter_cache", _spend_counter_cache(redis)) monkeypatch.setattr(ps, "prisma_client", prisma) reservation = _reservation(reserved_cost=0.4) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py index 54e5a6d5385..3f434978c36 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py @@ -34,6 +34,7 @@ async def test_spend_query_uses_timestamp_filtering(): mock_query_raw = AsyncMock(return_value=[]) mock_db.query_raw = mock_query_raw mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Use timezone-aware datetime objects start_date = datetime.datetime(2024, 1, 1, tzinfo=timezone.utc) @@ -87,6 +88,7 @@ async def test_global_activity_wraps_params_in_at_time_zone_utc(monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -134,6 +136,7 @@ async def test_global_activity_internal_user_wraps_params_in_at_time_zone_utc( mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -168,6 +171,7 @@ async def test_spend_logs_ui_wraps_params_in_at_time_zone_utc(monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_spendlogs = MagicMock() mock_prisma.db.litellm_spendlogs.count = AsyncMock(return_value=0) @@ -209,6 +213,7 @@ def _make_ui_spend_logs_mock(count_total, page_rows): """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( side_effect=[[{"total_count": count_total}], page_rows] ) @@ -259,6 +264,7 @@ async def test_spend_logs_ui_uses_bounded_count_not_full_scan(monkeypatch): ) mock_prisma.db.litellm_spendlogs.count.assert_not_called() + mock_prisma.replica_db = mock_prisma.db count_call = mock_prisma.db.query_raw.call_args_list[0] count_sql = count_call[0][0] @@ -347,6 +353,7 @@ async def test_spend_logs_ui_empty_page_reports_zero_total(monkeypatch): # page. mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(side_effect=[[{"total_count": 0}], []]) mock_prisma.db.litellm_spendlogs = MagicMock() mock_prisma.db.litellm_spendlogs.count = AsyncMock(return_value=0) @@ -395,6 +402,7 @@ async def test_spend_logs_ui_out_of_range_page_keeps_total(monkeypatch): # out-of-range page (empty). mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(side_effect=[[{"total_count": 7}], []]) mock_prisma.db.litellm_spendlogs = MagicMock() mock_prisma.db.litellm_spendlogs.count = AsyncMock(return_value=0) @@ -436,6 +444,7 @@ async def test_get_spend_by_team_binds_optional_team_filter(): """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_query_raw = AsyncMock(return_value=[]) mock_prisma.db.query_raw = mock_query_raw @@ -487,6 +496,7 @@ async def test_global_spend_report_team_group_forwards_team_id(monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -545,6 +555,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch): mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(side_effect=mock_query_raw) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index b3913079bb2..fdbe335d3b5 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -426,6 +426,7 @@ async def test_should_shrink_second_tag_reservation_to_remaining_budget( ) prisma_client = MagicMock() prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[]) + prisma_client.replica_db = prisma_client.db with patch( "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", @@ -2561,7 +2562,7 @@ async def test_reconcile_before_db_update_does_not_double_count_when_flush_lands counter_cache.redis_cache = redis_cache counter_cache.in_memory_cache.set_cache(key=counter_key, value=0.6) db_floor = _TeamMembershipFloorDb(spend=0.3) - ps.prisma_client = SimpleNamespace(db=db_floor) + ps.prisma_client = SimpleNamespace(db=db_floor, replica_db=db_floor) reservation = { "reserved_cost": 0.6, @@ -3476,6 +3477,7 @@ class _ModelAccessGroupBudgetPrisma: self.db = SimpleNamespace( litellm_modelaccessgroupbudgettable=SimpleNamespace(find_many=self._find_many) ) + self.replica_db = self.db async def _find_many(self, **kwargs): requested = list(kwargs["where"]["access_group_name"]["in"]) diff --git a/tests/test_litellm/proxy/test_fallback_management_endpoints.py b/tests/test_litellm/proxy/test_fallback_management_endpoints.py index 054dafbf5a7..4f4b9ed2650 100644 --- a/tests/test_litellm/proxy/test_fallback_management_endpoints.py +++ b/tests/test_litellm/proxy/test_fallback_management_endpoints.py @@ -123,6 +123,7 @@ class TestCreateFallback: """Create a mock prisma client""" client = MagicMock() client.db.litellm_config.upsert = AsyncMock() + client.replica_db = client.db client.jsonify_object = lambda x: x return client @@ -178,6 +179,7 @@ class TestCreateFallback: # Verify database was updated mock_prisma_client.db.litellm_config.upsert.assert_called_once() + mock_prisma_client.replica_db = mock_prisma_client.db async def test_create_fallback_router_not_initialized( self, mock_prisma_client, mock_proxy_config, mock_user_api_key_dict @@ -430,6 +432,7 @@ class TestDeleteFallback: """Create a mock prisma client""" client = MagicMock() client.db.litellm_config.upsert = AsyncMock() + client.replica_db = client.db client.jsonify_object = lambda x: x return client @@ -487,6 +490,7 @@ class TestDeleteFallback: # Verify database was updated mock_prisma_client.db.litellm_config.upsert.assert_called_once() + mock_prisma_client.replica_db = mock_prisma_client.db async def test_delete_fallback_not_found( self, diff --git a/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py index 1a514ed2c57..15fdf37dc11 100644 --- a/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py +++ b/tests/test_litellm/proxy/test_filter_models_by_team_access_group.py @@ -76,6 +76,7 @@ async def test_filter_resolves_access_group_names(): # Prisma mock mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) result = await _filter_models_by_team_id( @@ -129,6 +130,7 @@ async def test_filter_resolves_mix_of_access_groups_and_literal_names(): mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) result = await _filter_models_by_team_id( @@ -173,6 +175,7 @@ async def test_filter_excludes_models_from_other_access_group(): mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) result = await _filter_models_by_team_id( @@ -212,6 +215,7 @@ async def test_filter_db_fallback_receives_resolved_model_names(): mock_prisma = MagicMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock( return_value=[mock_db_model] ) diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index 7d6c5d3cebe..79c0292da3b 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -25,6 +25,7 @@ def mock_prisma(): """Simplified mock PrismaClient with bound methods""" client = MagicMock() client.db.litellm_healthchecktable.create = AsyncMock(return_value={"id": "test-id"}) + client.replica_db = client.db client.db.litellm_healthchecktable.find_many = AsyncMock(return_value=[{"id": "1", "model_name": "test"}]) # Bind actual methods @@ -55,6 +56,7 @@ async def test_save_health_check_result(mock_prisma, status, healthy, unhealthy, """Test health check result saving with various scenarios""" if not should_succeed: mock_prisma.db.litellm_healthchecktable.create.side_effect = Exception("DB Error") + mock_prisma.replica_db = mock_prisma.db result = await mock_prisma.save_health_check_result( model_name="test-model", @@ -74,6 +76,7 @@ async def test_get_health_check_history(mock_prisma): """Test health check history retrieval""" result = await mock_prisma.get_health_check_history(model_name="test", limit=50) mock_prisma.db.litellm_healthchecktable.find_many.assert_called_once() + mock_prisma.replica_db = mock_prisma.db assert len(result) == 1 @@ -375,6 +378,7 @@ async def test_save_background_health_checks_to_db(): mock_prisma = MagicMock() mock_prisma.save_health_check_result = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db model_list = [ { @@ -494,6 +498,7 @@ def _one_model_setup(): async def test_save_background_health_checks_to_db_returns_false_when_a_write_fails(): mock_prisma = MagicMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) + mock_prisma.replica_db = mock_prisma.db mock_prisma.save_health_check_result = AsyncMock(return_value=None) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() @@ -511,6 +516,7 @@ async def test_save_background_health_checks_to_db_writes_nothing_when_the_lates cycle by every pod while the read kept failing, which is what filled the table in production. """ mock_prisma.db.query_raw = AsyncMock(side_effect=RuntimeError("db down")) + mock_prisma.replica_db = mock_prisma.db mock_prisma.save_health_check_result = AsyncMock(return_value={"id": "row"}) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() @@ -533,6 +539,7 @@ async def test_save_background_health_checks_to_db_exception_handling(): """Test exception handling in background health check save""" mock_prisma = MagicMock() mock_prisma.db.query_raw = AsyncMock(side_effect=Exception("DB Error")) + mock_prisma.replica_db = mock_prisma.db model_list = [ { @@ -584,6 +591,7 @@ async def test_get_all_latest_health_checks_keeps_every_distinct_group_with_its_ _raw_latest_row("gpt-4", None, now - timedelta(minutes=3)), ] ) + mock_prisma.replica_db = mock_prisma.db result = await mock_prisma.get_all_latest_health_checks() @@ -609,6 +617,7 @@ async def test_save_background_health_checks_compares_raw_checked_at_against_utc _raw_latest_row("fresh-model", "fresh-id", fresh), ] ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.save_health_check_result = AsyncMock() model_list = [ {"model_name": "stale-model", "model_info": {"id": "stale-id"}, "litellm_params": {"model": "openai/stale"}}, diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 35308474949..6d109e089b8 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -149,6 +149,7 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re prisma_client.db.litellm_agentstable.find_unique = AsyncMock( side_effect=[None, _DbAgentRow("a2a-sibling-replica-agent-id", agent_name)] ) + prisma_client.replica_db = prisma_client.db monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index 0b51062dd66..45785fcee6b 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -1151,7 +1151,8 @@ def _fake_prisma_client_with_models(rows): from types import SimpleNamespace table = FakeProxyModelTable(rows) - return SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table)), table + tables = SimpleNamespace(litellm_proxymodeltable=table) + return SimpleNamespace(db=tables, replica_db=tables), table def _db_model_row(model_name: str, mock_response: str): @@ -1301,6 +1302,7 @@ async def test_route_request_a2a_agent_miss_does_not_consume_model_read_through( fake_prisma, model_table = _fake_prisma_client_with_models([]) agents_find_unique = AsyncMock(return_value=None) fake_prisma.db.litellm_agentstable = SimpleNamespace(find_unique=agents_find_unique) + fake_prisma.replica_db = fake_prisma.db monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma) monkeypatch.setattr(proxy_server, "store_model_in_db", True) monkeypatch.setattr(proxy_server, "llm_router", router) diff --git a/tests/test_litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py index ace4c4e65af..507e9c4c376 100644 --- a/tests/test_litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -65,6 +65,7 @@ def happy_path_upsert(monkeypatch): prisma_client = MagicMock() prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_teamtable.update = AsyncMock() class _FakeTx: diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 0140fcaba21..edef6100d4c 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -348,6 +348,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # Mock decryption to return the values as-is (simulating decryption) @@ -414,6 +415,7 @@ class TestProxySettingEndpoints: mock_db_record = MagicMock() mock_db_record.sso_settings = sso_settings mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_db_record) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # The resolver decrypts stored values via decrypt_value_helper; make it an @@ -554,6 +556,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) @@ -639,6 +642,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) @@ -704,6 +708,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock( @@ -762,6 +767,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() @@ -842,6 +848,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() env_var_entry = MagicMock() @@ -913,6 +920,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() @@ -991,6 +999,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) @@ -1336,6 +1345,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_uisettings.find_unique = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) response = client.get("/get/ui_settings") @@ -1368,6 +1378,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_uisettings.find_unique = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma) store = SettingsStore("general_settings") @@ -1427,6 +1438,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) response = client.get("/get/ui_settings") @@ -1463,6 +1475,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_uisettings.find_unique = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) class MockUser: @@ -1491,6 +1504,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_uisettings.find_unique.assert_called_once_with( where={"id": "ui_settings"} ) + mock_prisma.replica_db = mock_prisma.db def test_update_ui_settings_allowlisted_value(self, mock_auth, monkeypatch): """Test updating UI settings with an allowlisted field""" @@ -1509,6 +1523,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1551,6 +1566,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1595,6 +1611,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1632,6 +1649,7 @@ class TestProxySettingEndpoints: monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1678,6 +1696,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1715,6 +1734,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1752,6 +1772,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1798,6 +1819,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -1852,6 +1874,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() upsert_mock = AsyncMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = upsert_mock mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) @@ -1930,6 +1953,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() @@ -1982,6 +2006,7 @@ class TestProxySettingEndpoints: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() @@ -2026,6 +2051,7 @@ class TestProxySettingEndpoints: # Mock the prisma client to return None (no record found) mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -2111,6 +2137,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) # Mock decryption to return the values as-is (role_mappings should not be passed to decryption) @@ -2156,6 +2183,7 @@ class TestProxySettingEndpoints: # Mock the prisma client mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock() mock_prisma.db.litellm_config = MagicMock() mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) @@ -2298,6 +2326,7 @@ class TestProxySettingEndpoints: # Mock the prisma client to return None (no database record) mock_prisma = MagicMock() mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db # Run the async function role_mappings = asyncio.run(_setup_role_mappings()) @@ -2335,6 +2364,7 @@ class TestProxySettingEndpoints: mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock( return_value=mock_db_record ) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) from litellm.proxy.proxy_server import proxy_config @@ -2391,6 +2421,7 @@ def test_update_internal_user_settings_writes_audit_log(mock_proxy_config, monke audit_create = AsyncMock() fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = audit_create + fake_prisma.replica_db = fake_prisma.db monkeypatch.setattr(proxy_server_module, "prisma_client", fake_prisma) monkeypatch.setattr(proxy_server_module, "premium_user", True) @@ -2480,6 +2511,7 @@ def test_update_sso_settings_writes_redacted_audit_log(mock_proxy_config, monkey audit_create = AsyncMock() fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = audit_create + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_ssoconfig.upsert = AsyncMock() # No prior SSO row, so before_value resolves to None. fake_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) @@ -2545,6 +2577,7 @@ def test_update_sso_settings_audit_captures_redacted_before_snapshot( audit_create = AsyncMock() fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = audit_create + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_ssoconfig.upsert = AsyncMock() # Pre-existing SSO row contains the *prior* secret (would be ciphertext in @@ -2624,6 +2657,7 @@ def test_add_allowed_ip_writes_audit_log(mock_proxy_config, monkeypatch): audit_create = AsyncMock() fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = audit_create + fake_prisma.replica_db = fake_prisma.db monkeypatch.setattr(proxy_server_module, "prisma_client", fake_prisma) monkeypatch.setattr(proxy_server_module, "premium_user", True) @@ -2684,6 +2718,7 @@ def test_add_allowed_ip_hands_save_config_only_the_changed_general_setting(monke fake_prisma: Final = MagicMock() fake_prisma.db.litellm_auditlog.create = AsyncMock() + fake_prisma.replica_db = fake_prisma.db save_config: Final = AsyncMock(side_effect=lambda new_config: new_config) async def _get_config(): @@ -2731,6 +2766,7 @@ def test_delete_allowed_ip_writes_deleted_audit_log(monkeypatch): audit_create = AsyncMock() fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = audit_create + fake_prisma.replica_db = fake_prisma.db config = {"general_settings": {"allowed_ips": ["203.0.113.77", "198.51.100.1"]}} @@ -2791,6 +2827,7 @@ def test_allowed_ip_routes_refuse_a_config_owned_list_with_a_clear_400(route, mo fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = AsyncMock() + fake_prisma.replica_db = fake_prisma.db async def _get_config(): return {"general_settings": {"allowed_ips": ["203.0.113.77"]}} @@ -2837,6 +2874,7 @@ def test_update_ui_theme_settings_writes_audit_log(mock_proxy_config, monkeypatc audit_create = AsyncMock() fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = audit_create + fake_prisma.replica_db = fake_prisma.db monkeypatch.setattr(proxy_server_module, "prisma_client", fake_prisma) monkeypatch.setattr(proxy_server_module, "premium_user", True) @@ -2886,6 +2924,7 @@ def test_update_ui_settings_writes_audit_log(monkeypatch): audit_create = AsyncMock() fake_prisma = MagicMock() fake_prisma.db.litellm_auditlog.create = audit_create + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) fake_prisma.db.litellm_uisettings.upsert = AsyncMock() @@ -2940,6 +2979,7 @@ def mock_team_lookup(monkeypatch): find_many = AsyncMock(side_effect=_find_many) fake_prisma = MagicMock() fake_prisma.db.litellm_teamtable.find_many = find_many + fake_prisma.replica_db = fake_prisma.db member_budget_update = AsyncMock() @@ -3077,6 +3117,7 @@ def mock_organization_lookup(monkeypatch): find_unique = AsyncMock(side_effect=_find_unique) fake_prisma = MagicMock() fake_prisma.db.litellm_organizationtable.find_unique = find_unique + fake_prisma.replica_db = fake_prisma.db monkeypatch.setattr(proxy_server_module, "prisma_client", fake_prisma) monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) @@ -3516,6 +3557,7 @@ class TestPtuCostAttributionUISetting: mock_record = MagicMock() mock_record.ui_settings = stored mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=mock_record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.upsert = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) return mock_prisma @@ -3700,6 +3742,7 @@ class TestApplyUserBudgetToTeamKeysUISetting: mock_record = MagicMock() mock_record.ui_settings = stored mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=mock_record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.upsert = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) return mock_prisma @@ -3782,6 +3825,7 @@ class TestTeamAdminEditableTeamFieldsSetting: monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.upsert = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) return mock_prisma @@ -3876,6 +3920,7 @@ class TestTeamAdminEditableTeamFieldsSetting: mock_db_record = MagicMock() mock_db_record.ui_settings = {"team_admin_editable_team_fields": ["tpm_limit"]} mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=mock_db_record) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) general_settings: dict = {} monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings) @@ -3904,6 +3949,7 @@ class TestTeamAdminEditableTeamFieldsSetting: mock_db_record = MagicMock() mock_db_record.ui_settings = {"team_admin_editable_team_fields": stored} mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=mock_db_record) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) try: @@ -3944,6 +3990,7 @@ class TestSyncUiSettingsToGeneralSettings: } ) mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=record) + mock_prisma.replica_db = mock_prisma.db applied = await self._sync()(mock_prisma) @@ -3965,6 +4012,7 @@ class TestSyncUiSettingsToGeneralSettings: record = MagicMock() record.ui_settings = {"team_admin_editable_team_fields": ["rpm_limit"]} mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=record) + mock_prisma.replica_db = mock_prisma.db await self._sync()(mock_prisma) @@ -3978,6 +4026,7 @@ class TestSyncUiSettingsToGeneralSettings: monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", general_settings) mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None) + mock_prisma.replica_db = mock_prisma.db applied = await self._sync()(mock_prisma) diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_user_banner_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_user_banner_endpoints.py index 5d4875073fd..a81f2bb515a 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_user_banner_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_user_banner_endpoints.py @@ -52,6 +52,7 @@ def mock_audit_log(monkeypatch): def _mock_prisma(monkeypatch, record=None): mock_prisma = MagicMock() mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=record) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_uisettings.upsert = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) return mock_prisma @@ -105,6 +106,7 @@ class TestUpdateUserBanner: response = client.patch("/update/user_banner", json=PUBLISH_BODY) assert response.status_code == 403 mock_prisma.db.litellm_uisettings.upsert.assert_not_awaited() + mock_prisma.replica_db = mock_prisma.db def test_persists_and_round_trips(self, admin_auth, monkeypatch, mock_audit_log): mock_prisma = _mock_prisma(monkeypatch, record=None) @@ -124,6 +126,7 @@ class TestUpdateUserBanner: mock_prisma.db.litellm_uisettings.find_unique = AsyncMock( return_value=SimpleNamespace(ui_settings=persisted_payload) ) + mock_prisma.replica_db = mock_prisma.db read_back = client.get("/get/user_banner") assert read_back.status_code == 200 assert read_back.json() == saved @@ -173,3 +176,4 @@ class TestUpdateUserBanner: response = client.patch("/update/user_banner", json=payload) assert response.status_code == 422 mock_prisma.db.litellm_uisettings.upsert.assert_not_awaited() + mock_prisma.replica_db = mock_prisma.db diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py index c502fe4800e..8339fcae6e7 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py @@ -122,6 +122,7 @@ def mock_prisma_client() -> MagicMock: """ client = MagicMock(name="MockPrismaClient") client.db = MagicMock(name="MockPrismaDB") + client.replica_db = client.db client.connect = AsyncMock() client.disconnect = AsyncMock() client.health_check = AsyncMock(return_value=[{"?column?": 1}]) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py index 761835078f4..70ac637ec76 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_config_param_cache.py @@ -233,6 +233,7 @@ async def test_prefetch_config_params_populates_cache_for_each_name( ] prisma = MagicMock() prisma.db.litellm_config.find_many = AsyncMock(return_value=rows) + prisma.replica_db = prisma.db await prefetch_config_params(prisma, ["a", "b", "c"]) actual = { "a": _swap_config_cache._store[_config_cache_key("a")], @@ -252,6 +253,7 @@ async def test_prefetch_config_params_empty_list_is_noop( ) -> None: prisma = MagicMock() prisma.db.litellm_config.find_many = AsyncMock(return_value=[]) + prisma.replica_db = prisma.db await prefetch_config_params(prisma, []) assert prisma.db.litellm_config.find_many.await_count == 0 assert _swap_config_cache._store == {} @@ -263,5 +265,6 @@ async def test_prefetch_config_params_swallows_db_error_without_caching( ) -> None: prisma = MagicMock() prisma.db.litellm_config.find_many = AsyncMock(side_effect=RuntimeError("boom")) + prisma.replica_db = prisma.db await prefetch_config_params(prisma, ["a", "b"]) assert _swap_config_cache._store == {} diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py index 3c028473479..8e43ecda31f 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_password_helpers.py @@ -147,6 +147,7 @@ def _make_user(user_id: str, password) -> SimpleNamespace: async def test_migrate_passwords_skips_when_no_plaintext() -> None: pc = MagicMock() pc.db = MagicMock() + pc.replica_db = pc.db sha = hashlib.sha256(b"already-hashed").hexdigest() pc.db.litellm_usertable.find_many = AsyncMock( return_value=[ @@ -175,6 +176,7 @@ async def test_migrate_passwords_skips_when_no_plaintext() -> None: async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: pc = MagicMock() pc.db = MagicMock() + pc.replica_db = pc.db users: List[SimpleNamespace] = [ _make_user("plaintext-user-1", "plain-1"), _make_user("plaintext-user-2", "plain-2"), @@ -216,6 +218,7 @@ async def test_migrate_passwords_upgrades_only_plaintext_rows() -> None: async def test_migrate_passwords_raises_on_db_failure() -> None: pc = MagicMock() pc.db = MagicMock() + pc.replica_db = pc.db pc.db.litellm_usertable.find_many = AsyncMock( side_effect=RuntimeError("db unavailable") ) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py index 7099101db1c..18102ea2217 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py @@ -48,6 +48,7 @@ async def test_update_end_user_spend_upserts_each_end_user( transaction = MagicMock() transaction.batch_ = lambda: _AsyncCM(batcher) mock_prisma_client.db.tx = lambda timeout: _AsyncCM(transaction) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -98,6 +99,7 @@ async def test_update_end_user_spend_retries_on_connect_error( err = httpx.ConnectError("down") mock_prisma_client.db.tx = MagicMock(side_effect=err) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() with pytest.raises(httpx.ConnectError): @@ -122,6 +124,7 @@ async def test_update_end_user_spend_does_not_retry_post_send_ambiguous_errors( err = getattr(httpx, ambiguous_error_name)("ambiguous") mock_prisma_client.db.tx = MagicMock(side_effect=err) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() with pytest.raises((httpx.ReadTimeout, httpx.ReadError)): @@ -139,6 +142,7 @@ async def test_update_end_user_spend_non_connection_error_raises_immediately( mock_prisma_client: Any, ) -> None: mock_prisma_client.db.tx = MagicMock(side_effect=RuntimeError("unknown")) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() with pytest.raises(RuntimeError, match="unknown"): @@ -182,6 +186,7 @@ async def test_update_end_user_spend_retries_on_deadlock_then_commits( transaction = MagicMock() transaction.batch_ = lambda: _AsyncCM(batcher) mock_prisma_client.db.tx = MagicMock(side_effect=[_failing_tx(_end_user_deadlock_error()), _AsyncCM(transaction)]) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -209,6 +214,7 @@ async def test_update_end_user_spend_raises_after_exhausting_deadlock_retries( monkeypatch.setattr(asyncio, "sleep", AsyncMock(return_value=None)) mock_prisma_client.db.tx = MagicMock(side_effect=lambda timeout: _failing_tx(_end_user_deadlock_error())) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -229,6 +235,7 @@ async def test_update_spend_logs_writes_batches_via_create_many( ) -> None: logs = [make_spend_log_row(request_id=f"r{i}", spend=float(i)) for i in range(3)] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() await ProxyUpdateSpend.update_spend_logs( @@ -265,6 +272,7 @@ async def test_update_spend_logs_bounds_each_statement_by_payload_bytes( blob = json.dumps({"content": "x" * 10_000}) logs = [make_spend_log_row(request_id=f"r{i}", messages=blob, response=blob) for i in range(50)] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -329,6 +337,7 @@ async def test_update_spend_logs_pops_logs_when_logs_to_process_is_none( make_spend_log_row(request_id="b"), ] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() await ProxyUpdateSpend.update_spend_logs( @@ -359,6 +368,7 @@ async def test_update_spend_logs_failure_raises_after_retries( monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=httpx.ReadError("network blip")) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() with pytest.raises(httpx.ReadError): @@ -397,6 +407,7 @@ async def test_update_spend_logs_isolates_poison_row_and_persists_good_rows( written.extend(ids) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_create_many) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -421,6 +432,7 @@ async def test_update_spend_logs_reraises_connection_masquerade_dataerror( """ err = _data_error("Can't reach database server at db-host:5432") mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=err) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -454,6 +466,7 @@ async def test_update_spend_logs_retries_and_requeues_batch_on_db_outage( mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( side_effect=_data_error("Can't reach database server at db-host:5432 (P1001)") ) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() logs = [make_spend_log_row(request_id="a"), make_spend_log_row(request_id="b")] @@ -496,6 +509,7 @@ async def test_update_spend_logs_retries_deadlock_and_keeps_every_row( monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) create_many = AsyncMock(side_effect=[_deadlock_error(), _deadlock_error(), None]) mock_prisma_client.db.litellm_spendlogs.create_many = create_many + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() mock_prisma_client.spend_log_transactions = [] @@ -528,6 +542,7 @@ async def test_update_spend_logs_requeues_batch_once_deadlock_retries_exhaust( monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_deadlock_error()) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="c")] @@ -560,6 +575,7 @@ async def test_update_spend_logs_requeues_batch_on_non_transport_db_error( {"user_facing_error": {"error_code": "P2021", "message": "The table does not exist", "meta": {}}} ) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=err) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="c")] @@ -658,6 +674,7 @@ async def test_update_spend_logs_does_not_requeue_non_transport_failures( wedge the head of the queue forever. """ mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=ValueError("bad payload")) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() mock_prisma_client.spend_log_transactions = [] @@ -699,6 +716,7 @@ async def test_update_spend_logs_caps_isolation_attempts_under_poison_flood( raise _data_error("invalid byte sequence for encoding UTF8: 0x00") mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_always_poison) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() logs = [make_spend_log_row(request_id=f"r{i}") for i in range(n_rows)] @@ -763,6 +781,7 @@ async def _flush_and_count_create_many( monkeypatch.setattr(utils_mod, "SPEND_LOG_WRITE_BATCH_MAX_BYTES", max_bytes) monkeypatch.setattr(utils_mod, "SPEND_LOG_WRITE_BATCH_MAX_ROWS", max_rows) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_create_many) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -831,6 +850,7 @@ async def test_clean_statement_is_still_written_after_a_poison_flood( monkeypatch.setattr(utils_mod, "SPEND_LOG_WRITE_BATCH_MAX_BYTES", 250_000) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_create_many) + mock_prisma_client.replica_db = mock_prisma_client.db proxy_logging = MagicMock() proxy_logging.failure_handler = AsyncMock() @@ -902,6 +922,7 @@ async def test_update_spend_logs_parks_failed_batch_in_redis_with_wire_safe_date {"user_facing_error": {"error_code": "P2021", "message": "The table does not exist", "meta": {}}} ) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=err) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.spend_log_transactions = [] with pytest.raises(TableNotFoundError): diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py index d6f41ba55db..ee59be84db7 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -186,6 +186,7 @@ async def test_update_spend_logs_job_skips_when_queue_empty( proxy_logging.failure_handler = AsyncMock() mock_prisma_client.spend_log_transactions = [] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db await update_spend_logs_job( prisma_client=mock_prisma_client, db_writer_client=None, @@ -209,6 +210,7 @@ async def test_update_spend_logs_job_drains_tool_queue_when_spend_queue_empty( mock_prisma_client.spend_log_transactions = [] mock_prisma_client.tool_usage_transactions = [MagicMock()] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False) flush_stub = AsyncMock() monkeypatch.setattr(tool_mod, "flush_tool_usage_transactions", flush_stub, raising=False) @@ -234,6 +236,7 @@ async def test_update_spend_logs_job_processes_and_clears_queue( make_spend_log_row(request_id="r2"), ] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Stub auxiliary imports so the test focuses on the spend-logs write path. import litellm.proxy.guardrails.usage_tracking as guard_mod @@ -289,6 +292,7 @@ async def test_update_spend_logs_job_requeues_popped_rows_when_write_cancelled( mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( side_effect=_cancel_mid_write ) + mock_prisma_client.replica_db = mock_prisma_client.db with pytest.raises(asyncio.CancelledError): await update_spend_logs_job( @@ -315,6 +319,7 @@ async def test_update_spend_logs_job_does_not_requeue_when_cancelled_after_write proxy_logging.failure_handler = AsyncMock() mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr( guard_mod, @@ -361,6 +366,7 @@ async def test_drain_spend_logs_queue_flushes_rows_queued_while_draining( ) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write) + mock_prisma_client.replica_db = mock_prisma_client.db await drain_spend_logs_queue( prisma_client=mock_prisma_client, @@ -402,6 +408,7 @@ async def test_drain_spend_logs_queue_stops_monitor_and_keeps_its_popped_rows( written.extend(row["request_id"] for row in kwargs["data"]) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write) + mock_prisma_client.replica_db = mock_prisma_client.db async def _monitor() -> None: await update_spend_logs_job( @@ -448,6 +455,7 @@ async def test_drain_spend_logs_queue_gives_up_after_max_passes( mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( side_effect=_write_and_refill ) + mock_prisma_client.replica_db = mock_prisma_client.db await drain_spend_logs_queue( prisma_client=mock_prisma_client, @@ -747,6 +755,7 @@ async def test_drain_spend_logs_queue_parks_unwritable_rows_in_redis_on_shutdown make_spend_log_row(request_id="r2"), ] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_table_gone_error()) + mock_prisma_client.replica_db = mock_prisma_client.db with pytest.raises(TableNotFoundError): await drain_spend_logs_queue( @@ -771,6 +780,7 @@ async def test_drain_spend_logs_queue_waits_for_an_in_flight_write_before_parkin raise _table_gone_error() mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_fail_once_shutdown_starts) + mock_prisma_client.replica_db = mock_prisma_client.db scheduler_write: Final = asyncio.ensure_future( update_spend_logs_job( prisma_client=mock_prisma_client, @@ -818,6 +828,7 @@ async def test_drain_spend_logs_queue_parks_rows_left_after_max_passes( mock_prisma_client.spend_log_transactions.append(make_spend_log_row(request_id="late")) mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write_and_refill) + mock_prisma_client.replica_db = mock_prisma_client.db await drain_spend_logs_queue( prisma_client=mock_prisma_client, @@ -838,6 +849,7 @@ async def test_drain_spend_logs_queue_keeps_rows_in_memory_when_redis_is_down( fake_redis.down = True mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_table_gone_error()) + mock_prisma_client.replica_db = mock_prisma_client.db with pytest.raises(TableNotFoundError): await drain_spend_logs_queue( @@ -867,6 +879,7 @@ async def test_update_spend_writes_rows_parked_in_redis_by_a_previous_pod( assert await buffer.store_spend_logs_in_redis([make_spend_log_row(request_id="parked")]) is True mock_prisma_client.spend_log_transactions = [] mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db await update_spend( prisma_client=mock_prisma_client, diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py index dc4f2900038..87c32500007 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py @@ -121,6 +121,7 @@ async def test_delete_vector_store_checks_access(): } ) mock_prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=mock_vector_store) + mock_prisma.replica_db = mock_prisma.db # User from different team should get 403 user_api_key_dict = UserAPIKeyAuth(team_id="team_789") @@ -279,6 +280,7 @@ async def test_get_vector_store_info_dashboard_session_resolves_real_teams( mock_prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=MagicMock(model_dump=lambda: dict(_TEAM_A_OWNED)) ) + mock_prisma.replica_db = mock_prisma.db async def outcome() -> int: try: diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 2f7d4b350be..87e4d3d784b 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -1428,6 +1428,7 @@ class TestIndexCreate: mock_prisma.db.litellm_managedvectorstoreindextable.find_unique = AsyncMock( return_value=None ) + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_managedvectorstoreindextable.create = AsyncMock( return_value=mock_row ) @@ -1481,6 +1482,7 @@ class TestIndexList: """Index topology must never reach non-admins, not even via a DB read.""" mock_prisma = MagicMock() mock_prisma.db.litellm_managedvectorstoreindextable.find_many = AsyncMock() + mock_prisma.replica_db = mock_prisma.db with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): with pytest.raises(HTTPException) as exc_info: @@ -1514,6 +1516,7 @@ class TestIndexList: ] mock_prisma = MagicMock() mock_prisma.db.litellm_managedvectorstoreindextable.find_many = AsyncMock(return_value=rows) + mock_prisma.replica_db = mock_prisma.db with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): result = await index_list(user_api_key_dict=self._admin()) @@ -1774,6 +1777,7 @@ async def test_vector_store_synchronization_across_instances(): mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( side_effect=mock_find_unique ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock( side_effect=mock_find_many ) @@ -2021,6 +2025,7 @@ async def test_vector_store_update_and_list_synchronization(): mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock( side_effect=mock_find_many ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock( side_effect=mock_create ) @@ -2170,6 +2175,7 @@ async def test_new_vector_store_persists_embedding_reference_without_credentials mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=None # Vector store doesn't exist yet ) + mock_prisma_client.replica_db = mock_prisma_client.db # Track what was passed to create captured_create_data = {} @@ -2186,6 +2192,7 @@ async def test_new_vector_store_persists_embedding_reference_without_credentials mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock( side_effect=mock_create ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_registry = MagicMock() mock_registry.add_vector_store_to_registry = MagicMock() @@ -2250,6 +2257,7 @@ async def test_new_vector_store_auto_resolves_from_router(): mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=None # Vector store doesn't exist yet ) + mock_prisma_client.replica_db = mock_prisma_client.db # Track what was passed to create captured_create_data = {} @@ -2265,6 +2273,7 @@ async def test_new_vector_store_auto_resolves_from_router(): return mock_created_vector_store mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock(side_effect=mock_create) + mock_prisma_client.replica_db = mock_prisma_client.db mock_registry = MagicMock() mock_registry.add_vector_store_to_registry = MagicMock() @@ -2387,6 +2396,7 @@ async def test_create_vector_store_in_db(): mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=None # Vector store doesn't exist yet ) + mock_prisma_client.replica_db = mock_prisma_client.db created_vector_store_data = { "vector_store_id": vector_store_id, @@ -2463,6 +2473,7 @@ async def test_create_vector_store_in_db_raises_when_exists(): mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=existing_vector_store ) + mock_prisma_client.replica_db = mock_prisma_client.db with pytest.raises(HTTPException) as exc_info: await create_vector_store_in_db( @@ -2699,6 +2710,7 @@ class TestUpdateVectorStoreAccessControlAndRedaction: mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=existing_row ) + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch( @@ -2768,6 +2780,7 @@ class TestUpdateVectorStoreAccessControlAndRedaction: mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=existing_row ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock( return_value=updated_row ) @@ -2819,6 +2832,7 @@ class TestUpdateVectorStoreAccessControlAndRedaction: mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( return_value=existing_row ) + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock( return_value=None ) @@ -3109,6 +3123,7 @@ class TestConfigOwnedVectorStores: registry = self._registry() prisma = MagicMock() prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[self._db_row(self.DB_ID, "db-store")]) + prisma.replica_db = prisma.db with ( patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam @@ -3130,6 +3145,7 @@ class TestConfigOwnedVectorStores: prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock( return_value=[self._db_row(self.DB_ID, "db-store"), self._db_row(self.CONFIG_ID, "renamed-in-db")] ) + prisma.replica_db = prisma.db with ( patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: proxy_server global, no seam @@ -3167,6 +3183,7 @@ class TestConfigOwnedVectorStores: async def test_new_with_config_store_id_is_rejected_before_db_write(self): prisma = MagicMock() prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + prisma.replica_db = prisma.db prisma.db.litellm_managedvectorstorestable.create = AsyncMock() with ( @@ -3191,6 +3208,7 @@ class TestConfigOwnedVectorStores: prisma = MagicMock() prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + prisma.replica_db = prisma.db prisma.db.litellm_managedvectorstorestable.update = AsyncMock() registry = self._registry() @@ -3216,6 +3234,7 @@ class TestConfigOwnedVectorStores: prisma = MagicMock() prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + prisma.replica_db = prisma.db prisma.db.litellm_managedvectorstorestable.delete = AsyncMock() registry = self._registry() @@ -3242,6 +3261,7 @@ class TestConfigOwnedVectorStores: row.model_dump = MagicMock(return_value=self._db_row(self.DB_ID, "db-store")) prisma = MagicMock() prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=row) + prisma.replica_db = prisma.db prisma.db.litellm_managedvectorstorestable.delete = AsyncMock() registry = self._registry() diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py index d717c4e8c89..4b0795236eb 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py @@ -315,6 +315,7 @@ async def test_load_state_from_db_adds_the_persisted_delta_to_the_cold_start_pri prisma = MagicMock() prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[fake_row]) + prisma.replica_db = prisma.db await r.load_state_from_db(prisma) new_cell = r._cells[(RequestType.GENERAL, "fast")] @@ -337,6 +338,7 @@ async def test_load_state_from_db_keeps_a_one_sided_delta_row_sampleable(): prisma = MagicMock() prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[one_sided_row]) + prisma.replica_db = prisma.db await r.load_state_from_db(prisma) loaded_cell = r._cells[(RequestType.GENERAL, "fast")] @@ -365,6 +367,7 @@ async def test_load_state_from_db_handles_unknown_request_type(): prisma = MagicMock() prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[bad_row, good_row]) + prisma.replica_db = prisma.db await r.load_state_from_db(prisma) # Unknown skipped; good added to the cold-start prior. diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py index 23fc859d4a6..df9679b780e 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py @@ -55,6 +55,7 @@ def _make_router( def _make_mock_prisma(): p = MagicMock() p.db.litellm_adaptiverouterstate.find_unique = AsyncMock(return_value=None) + p.replica_db = p.db p.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[]) p.db.litellm_adaptiverouterstate.upsert = AsyncMock() p.db.litellm_adaptiveroutersession.upsert = AsyncMock() @@ -200,6 +201,7 @@ async def test_load_state_from_db_adds_persisted_delta_to_cold_start(): prisma = _make_mock_prisma() prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[fake_row]) + prisma.replica_db = prisma.db await router.load_state_from_db(prisma) @@ -219,6 +221,7 @@ async def test_load_state_from_db_handles_unknown_request_type(): prisma = _make_mock_prisma() prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[bad_row]) + prisma.replica_db = prisma.db # Should not raise; bad row is silently skipped and cold-start cells remain. await router.load_state_from_db(prisma) diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_update_queue.py b/tests/test_litellm/router_strategy/adaptive_router/test_update_queue.py index 9baa69a19e0..120707d95cf 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_update_queue.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_update_queue.py @@ -18,6 +18,7 @@ def mock_prisma(): """Prisma client with both adaptive router models stubbed as AsyncMocks.""" p = MagicMock() p.db.litellm_adaptiverouterstate.find_unique = AsyncMock(return_value=None) + p.replica_db = p.db p.db.litellm_adaptiverouterstate.upsert = AsyncMock() p.db.litellm_adaptiveroutersession.upsert = AsyncMock() return p @@ -53,6 +54,7 @@ async def test_add_session_state_last_write_wins(queue): flushed.append(kwargs) p.db.litellm_adaptiveroutersession.upsert = upsert + p.replica_db = p.db await queue.flush_session_to_db(p) assert len(flushed) == 1 assert flushed[0]["data"]["update"]["misalignment_count"] == 5 diff --git a/tests/test_litellm/test_model_block_unblock.py b/tests/test_litellm/test_model_block_unblock.py index da63ed4a95a..542f7eb325c 100644 --- a/tests/test_litellm/test_model_block_unblock.py +++ b/tests/test_litellm/test_model_block_unblock.py @@ -33,6 +33,7 @@ def _setup_model_block_mocks(monkeypatch, *, updated_blocked: bool): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_proxymodeltable = model_table + mock_prisma_client.replica_db = mock_prisma_client.db mock_router = MagicMock() mock_router.get_model_ids.return_value = [model_id] diff --git a/tests/test_litellm/test_router_retry_policy_update.py b/tests/test_litellm/test_router_retry_policy_update.py index 0a3dcba325a..345b4c73eae 100644 --- a/tests/test_litellm/test_router_retry_policy_update.py +++ b/tests/test_litellm/test_router_retry_policy_update.py @@ -357,6 +357,7 @@ async def test_config_update_persists_and_reads_back_retry_policy(monkeypatch): fake_table = _FakeConfigTable() prisma_client = MagicMock() prisma_client.db.litellm_config = fake_table + prisma_client.replica_db = prisma_client.db async def _apply_router_settings(*args, **kwargs): await proxy_server.proxy_config._add_router_settings_from_db_config( diff --git a/tests/test_litellm/vector_stores/test_vector_store_registry.py b/tests/test_litellm/vector_stores/test_vector_store_registry.py index 762176d6a81..ce191bfdc53 100644 --- a/tests/test_litellm/vector_stores/test_vector_store_registry.py +++ b/tests/test_litellm/vector_stores/test_vector_store_registry.py @@ -238,6 +238,7 @@ async def test_config_owned_store_survives_db_liveness_check_while_missing_db_st registry.add_vector_store_to_registry(_db_store("vs_from_db", "db-store")) prisma_client = MagicMock() prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + prisma_client.replica_db = prisma_client.db to_run = await registry.pop_vector_stores_to_run_with_db_fallback( non_default_params={"vector_store_ids": ["vs_from_config", "vs_from_db"]},