From 74d9fe2ba0a273f2e9d9fa187c199b1b053539e2 Mon Sep 17 00:00:00 2001 From: yuneng Date: Thu, 24 Sep 2026 09:07:09 +0000 Subject: [PATCH] test: alias replica_db on remaining prisma test doubles Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_prometheus_logging_callbacks.py | 1 + .../proxy/hooks/test_managed_files.py | 88 ++++++++++++++++--- .../test_internal_user_endpoints.py | 4 + .../test_project_endpoints_prisma.py | 5 ++ .../proxy/test_audit_logging_endpoints.py | 1 + .../test_proxy_budget_reset.py | 15 ++++ tests/logging_callback_tests/test_alerting.py | 1 + .../test_openai_batches_endpoint.py | 4 + .../test_claude_code_marketplace.py | 1 + .../test_passthrough_managed_ids.py | 2 + .../litellm_proxy/test_skills_ownership.py | 30 ++++--- tests/unit/repositories/test_repositories.py | 10 ++- 12 files changed, 135 insertions(+), 27 deletions(-) diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 58cde4c8103..22dde8a4d2e 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -1740,6 +1740,7 @@ async def test_initialize_remaining_budget_metrics_exception_handling( mock_db.litellm_teamtable = mock_teamtable mock_db.litellm_organizationtable = mock_orgtable mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db # Mock the Prometheus metrics prometheus_logger.litellm_remaining_team_budget_metric = MagicMock() diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 7e96c956664..36661389bb3 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -16,9 +16,14 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( ) +def _prisma_double(client: MagicMock) -> MagicMock: + client.replica_db = client.db + return client + + def test_get_file_ids_from_messages(): proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) messages = [ { @@ -42,7 +47,7 @@ def test_get_file_ids_from_messages(): def test_get_file_ids_from_messages_skips_bedrock_content_blocks_without_type(): proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) messages = [ { @@ -78,6 +83,7 @@ async def test_async_pre_call_hook_batch_retrieve(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value @@ -107,6 +113,7 @@ async def test_list_user_batches_limit_zero_returns_empty_page_without_db_query( from litellm.proxy._types import UserAPIKeyAuth prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db proxy_managed_files = _PROXY_LiteLLMManagedFiles(DualCache(), prisma_client=prisma_client) page = await proxy_managed_files.list_user_batches( @@ -127,7 +134,7 @@ async def test_async_pre_call_deployment_hook_resolves_model_id_from_litellm_met file ID is resolved to the provider-specific file ID. """ proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) managed_file_id = "managed-file-abc" @@ -161,7 +168,7 @@ async def test_async_pre_call_deployment_hook_prefers_top_level_model_info(): should use it without falling back to litellm_metadata. """ proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) managed_file_id = "managed-file-abc" @@ -200,7 +207,7 @@ async def test_async_pre_call_deployment_hook_no_model_info_leaves_file_id_uncha the managed file ID should remain unchanged. """ proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) managed_file_id = "managed-file-abc" @@ -352,7 +359,7 @@ async def test_async_post_call_success_hook_for_unified_finetuning_job(): "model_id": "gpt-3.5-turbo-0613", } proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=AsyncMock() + DualCache(), prisma_client=_prisma_double(AsyncMock()) ) data = { "user_api_key_dict": {"parent_otel_span": MagicMock()}, @@ -373,6 +380,7 @@ async def test_async_pre_call_hook_for_unified_finetuning_job(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value @@ -405,6 +413,7 @@ async def test_can_user_call_unified_file_id(call_type): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedfiletable.find_first.return_value = return_value @@ -432,6 +441,7 @@ async def test_router_acreate_batch_only_selects_from_file_id_mapping(monkeypatc import litellm prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db return_value = MagicMock() return_value.created_by = "123" prisma_client.db.litellm_managedobjecttable.find_first.return_value = return_value @@ -523,7 +533,7 @@ async def test_output_file_id_for_batch_retrieve(): "unified_batch_id": "litellm_proxy;model_id:12345679;llm_batch_id:batch_685c5e5d63988190b85bdb2147ba131d", } proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=AsyncMock() + DualCache(), prisma_client=_prisma_double(AsyncMock()) ) response = await proxy_managed_files.async_post_call_success_hook( @@ -583,7 +593,7 @@ async def test_output_file_id_preserves_target_model_names_when_model_name_missi } proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=AsyncMock() + DualCache(), prisma_client=_prisma_double(AsyncMock()) ) provider_output_file = OpenAIFileObject( @@ -659,7 +669,7 @@ async def test_error_file_id_for_failed_batch(): } proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=AsyncMock() + DualCache(), prisma_client=_prisma_double(AsyncMock()) ) # Create a proper OpenAIFileObject for the error file @@ -707,6 +717,7 @@ async def test_async_post_call_success_hook_twice_assert_no_unique_violation(): # Use AsyncMock instead of real database connection prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db batch = LiteLLMBatch( id="bGl0ZWxsbV9wcm94eTttb2RlbF9pZDoxMjM0NTY3OTtsbG1fYmF0Y2hfaWQ6YmF0Y2hfNjg1YzVlNWQ2Mzk4ODE5MGI4NWJkYjIxNDdiYTEzMWQ", @@ -1105,7 +1116,7 @@ def test_get_file_ids_from_responses_tools(): file IDs from the tools parameter. """ proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) tools = [ @@ -1128,7 +1139,7 @@ def test_get_file_ids_from_responses_tools_multiple_tools(): Test that get_file_ids_from_responses_tools handles multiple tools. """ proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) tools = [ @@ -1162,7 +1173,7 @@ def test_get_file_ids_from_responses_tools_empty(): Test that get_file_ids_from_responses_tools handles empty or None tools. """ proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) # Test with None @@ -1192,6 +1203,7 @@ async def test_check_file_ids_access_with_unified_file_ids(): # Mock the access check to return True prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db internal_usage_cache = MagicMock() proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1229,6 +1241,7 @@ async def test_check_file_ids_access_denied(): unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db internal_usage_cache = MagicMock() proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1266,6 +1279,7 @@ async def test_check_file_ids_access_with_regular_files_only(): regular_file_id_2 = "file-xyz789" prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db internal_usage_cache = MagicMock() proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1301,6 +1315,7 @@ async def test_completion_with_file_access_check(): unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) internal_usage_cache = MagicMock() @@ -1361,6 +1376,7 @@ async def test_responses_with_file_access_check(): unified_file_id_2 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsNzc3Nzc3Nzc7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1YWVo7bGxtX291dHB1dF9maWxlX21vZGVsX2lkLG1vZGVsXzEyMw" prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) internal_usage_cache = MagicMock() @@ -1425,6 +1441,7 @@ async def test_store_unified_file_id_with_none_file_object(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.upsert = AsyncMock( return_value=MagicMock() ) @@ -1460,6 +1477,7 @@ async def test_store_unified_file_id_updates_file_metadata_on_existing_row(): from litellm.types.llms.openai import OpenAIFileObject prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.upsert = AsyncMock( return_value=MagicMock() ) @@ -1525,6 +1543,7 @@ async def test_afile_delete_returns_provider_response_when_stored_file_object_no unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsdGVzdC1pZDt0YXJnZXRfbW9kZWxfbmFtZXMsZ3B0LTRvO2xsbV9vdXRwdXRfZmlsZV9pZCxmaWxlLXByb3ZpZGVyLXh5ejtsbG1fb3V0cHV0X2ZpbGVfbW9kZWxfaWQsbW9kZWwtMTIz" prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db db_record = MagicMock() db_record.model_mappings = '{"model-123": "file-provider-xyz"}' prisma_client.db.litellm_managedfiletable.find_first = AsyncMock( @@ -1586,6 +1605,7 @@ async def test_afile_retrieve_fetches_from_provider_when_file_object_none(): from litellm.types.llms.openai import OpenAIFileObject prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db internal_usage_cache = MagicMock() proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1640,6 +1660,7 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none() and no llm_router is provided to fetch from the provider. """ prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db internal_usage_cache = MagicMock() proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1674,6 +1695,7 @@ async def test_afile_retrieve_returns_stored_file_object_when_exists(): from litellm.types.llms.openai import OpenAIFileObject prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db internal_usage_cache = MagicMock() proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1711,6 +1733,7 @@ async def test_afile_retrieve_raises_error_for_non_managed_file(): in the managed files table. """ prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db internal_usage_cache = MagicMock() proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1737,6 +1760,7 @@ async def test_list_batches_from_managed_objects_table(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db batch_record_1 = MagicMock() batch_record_1.unified_object_id = "unified-batch-id-1" @@ -1802,6 +1826,7 @@ async def test_list_batches_from_managed_objects_table_empty_list(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -1882,6 +1907,7 @@ async def test_list_batches_registers_and_returns_unified_output_file_ids(): ).decode() prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_many.return_value = [ _terminal_batch_record( unified_batch_uid, raw_input_file_id, raw_output_file_id, raw_error_file_id @@ -1965,6 +1991,7 @@ async def test_list_batches_resolves_existing_managed_rows_without_minting(): ] prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_many.return_value = records existing_rows = [ @@ -2000,6 +2027,7 @@ async def test_list_batches_caps_page_size_at_100(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] proxy_managed_files = _PROXY_LiteLLMManagedFiles( @@ -2023,6 +2051,7 @@ async def test_list_batches_from_managed_objects_table_provider_filter_raises_ex from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client @@ -2049,6 +2078,7 @@ async def test_list_batches_from_managed_objects_table_target_model_name_filter_ from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client @@ -2075,6 +2105,7 @@ async def test_list_batches_from_managed_objects_table_filters_by_created_by(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Create batch for user1 batch_user1 = MagicMock() @@ -2155,6 +2186,7 @@ async def test_list_batches_pagination_uses_unified_object_id_cursor(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_first.return_value = MagicMock() prisma_client.db.litellm_managedobjecttable.find_many.return_value = [] @@ -2248,6 +2280,7 @@ async def test_list_batches_pagination_walks_all_pages_without_loops_or_gaps(): ) prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( side_effect=fake_find_many ) @@ -2369,6 +2402,7 @@ async def test_list_batches_pagination_stable_when_created_at_ties(): ) prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( side_effect=fake_find_many ) @@ -2447,6 +2481,7 @@ def _fake_managed_object_table(rows): ) prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( side_effect=find_many ) @@ -2486,6 +2521,7 @@ async def test_list_batches_rejects_unknown_after_cursor(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( return_value=None ) @@ -2563,6 +2599,7 @@ async def test_list_batches_rejects_after_cursor_owned_by_another_user(): return other_users_batch prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( side_effect=find_first ) @@ -2783,6 +2820,7 @@ async def test_user_b_cannot_retrieve_user_a_batch(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return User A as the creator batch_record = MagicMock() @@ -2820,6 +2858,7 @@ async def test_user_b_cannot_cancel_user_a_batch(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return User A as the creator batch_record = MagicMock() @@ -2860,6 +2899,7 @@ async def test_user_a_can_retrieve_own_batch(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return User A as the creator batch_record = MagicMock() @@ -2898,6 +2938,7 @@ async def test_user_b_cannot_retrieve_user_a_file(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return User A as the creator file_record = MagicMock() @@ -2935,6 +2976,7 @@ async def test_user_b_cannot_download_user_a_file_content(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return User A as the creator file_record = MagicMock() @@ -2972,6 +3014,7 @@ async def test_user_b_cannot_delete_user_a_file(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return User A as the creator file_record = MagicMock() @@ -3011,6 +3054,7 @@ async def test_user_a_can_retrieve_own_file(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return User A as the creator file_record = MagicMock() @@ -3061,6 +3105,7 @@ async def test_list_batches_only_returns_user_own_batches(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Create batches for User A batch_user_a = MagicMock() @@ -3114,6 +3159,7 @@ async def test_same_user_different_keys_can_access_batch(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db # Mock database to return the user_id as creator batch_record = MagicMock() @@ -3209,6 +3255,7 @@ async def test_team_b_cannot_access_team_a_provider_format_batch( from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) @@ -3250,6 +3297,7 @@ async def test_authorized_callers_can_access_provider_format_batch(caller_kwargs from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) @@ -3279,6 +3327,7 @@ async def test_provider_format_batch_without_ownership_row_stays_accessible(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_first.return_value = None proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client @@ -3305,6 +3354,7 @@ async def test_fine_tuning_provider_format_id_not_enforced(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) @@ -3342,6 +3392,7 @@ async def test_team_b_cannot_access_team_a_provider_format_file( from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) @@ -3370,6 +3421,7 @@ async def test_same_team_can_access_provider_format_file(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") ) @@ -3394,6 +3446,7 @@ async def test_provider_format_file_without_ownership_row_stays_accessible(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.find_first.return_value = None proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(), prisma_client=prisma_client @@ -3417,6 +3470,7 @@ async def test_post_call_batch_create_stores_ownership_row(batch_id): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client ) @@ -3451,6 +3505,7 @@ async def test_post_call_batch_sync_does_not_claim_ownership(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0 proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client @@ -3473,6 +3528,7 @@ async def test_post_call_batch_sync_updates_existing_row(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1 prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") @@ -3507,6 +3563,7 @@ async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row( from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1 prisma_client.db.litellm_managedobjecttable.find_first.return_value = ( _owned_record(created_by="user_a", team_id="team_a") @@ -3542,6 +3599,7 @@ async def test_post_call_batch_create_does_not_store_output_file_ownership(): from litellm.proxy._types import UserAPIKeyAuth prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db proxy_managed_files = _PROXY_LiteLLMManagedFiles( MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client ) @@ -3591,6 +3649,7 @@ async def test_file_list_cursors_are_scoped_to_the_caller(): ) prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.find_many.return_value = [] proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client @@ -3648,6 +3707,7 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): "status": "processed", } prisma_client = AsyncMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_managedfiletable.find_many.return_value = [managed_row] proxy_managed_files = _PROXY_LiteLLMManagedFiles( DualCache(), prisma_client=prisma_client @@ -3672,7 +3732,7 @@ async def test_list_user_batches_provider_filter_rejected_with_400(): from litellm.proxy._types import ProxyException, UserAPIKeyAuth proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) with pytest.raises(ProxyException) as exc: @@ -3692,7 +3752,7 @@ async def test_list_user_batches_target_model_names_filter_rejected_with_400(): from litellm.proxy._types import ProxyException, UserAPIKeyAuth proxy_managed_files = _PROXY_LiteLLMManagedFiles( - DualCache(), prisma_client=MagicMock() + DualCache(), prisma_client=_prisma_double(MagicMock()) ) with pytest.raises(ProxyException) as exc: diff --git a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py index fbfb6a99726..08e65ac4e52 100644 --- a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_internal_user_endpoints.py @@ -55,6 +55,7 @@ class TestAvailableEnterpriseUsers: # Mock database count mock_prisma.db.litellm_usertable.count = _user_count(5) mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=2) + mock_prisma.replica_db = mock_prisma.db # Override the dependency client.app.dependency_overrides[mock_user_api_key_auth] = lambda: { @@ -91,6 +92,7 @@ class TestAvailableEnterpriseUsers: ): mock_prisma.db.litellm_usertable.count = _user_count(5, deactivated=2) mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=2) + mock_prisma.replica_db = mock_prisma.db client.app.dependency_overrides[mock_user_api_key_auth] = lambda: { "user_id": "test_user" @@ -124,6 +126,7 @@ class TestAvailableEnterpriseUsers: # Mock database count mock_prisma.db.litellm_usertable.count = _user_count(3) mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=1) + mock_prisma.replica_db = mock_prisma.db # Override the dependency client.app.dependency_overrides[mock_user_api_key_auth] = lambda: { @@ -161,6 +164,7 @@ class TestAvailableEnterpriseUsers: # Mock database count higher than max_users to trigger the bug mock_prisma.db.litellm_usertable.count = _user_count(8) mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=3) + mock_prisma.replica_db = mock_prisma.db # Override the dependency client.app.dependency_overrides[mock_user_api_key_auth] = lambda: { diff --git a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py index 36878fa698c..d632eaf6af3 100644 --- a/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py +++ b/tests/enterprise/litellm_enterprise/proxy/management_endpoints/test_project_endpoints_prisma.py @@ -835,6 +835,7 @@ async def test_list_projects_returns_timestamps(): fake_project.updated_at = now mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_projecttable.find_many = AsyncMock( return_value=[fake_project] ) @@ -885,6 +886,7 @@ async def test_update_project_invalidates_cached_project_object(monkeypatch): stale_row.model_dump = lambda: {"project_id": project_id, "team_id": None, "models": []} mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.jsonify_object = lambda data: data mock_prisma.db.litellm_projecttable.find_unique = AsyncMock(return_value=stale_row) @@ -945,6 +947,7 @@ async def test_delete_project_invalidates_cached_project_object(monkeypatch): row.model_dump = lambda: {"project_id": project_id, "team_id": None, "models": ["gpt-5.5"]} mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_projecttable.find_unique = AsyncMock(return_value=row) seeded = await get_project_object( @@ -995,6 +998,7 @@ async def test_update_project_succeeds_when_cache_eviction_fails(monkeypatch): updated_row = MagicMock() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.jsonify_object = lambda data: data mock_prisma.db.litellm_projecttable.find_unique = AsyncMock(return_value=existing_row) mock_prisma.db.litellm_projecttable.update = AsyncMock(return_value=updated_row) @@ -1233,6 +1237,7 @@ def _project_update_mocks(monkeypatch, stored_metadata: dict) -> mock.MagicMock: mock_prisma.jsonify_object = lambda data: data mock_prisma.db.litellm_projecttable.find_unique = mock.AsyncMock(return_value=existing_row) mock_prisma.db.litellm_projecttable.update = mock.AsyncMock(return_value=mock.MagicMock()) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", True) monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma) diff --git a/tests/enterprise/litellm_enterprise/proxy/test_audit_logging_endpoints.py b/tests/enterprise/litellm_enterprise/proxy/test_audit_logging_endpoints.py index fd1b05ff060..ad8f113b74e 100644 --- a/tests/enterprise/litellm_enterprise/proxy/test_audit_logging_endpoints.py +++ b/tests/enterprise/litellm_enterprise/proxy/test_audit_logging_endpoints.py @@ -37,6 +37,7 @@ def mock_prisma_client(): mock.db.litellm_auditlog.find_many = AsyncMock() mock.db.litellm_auditlog.find_unique = AsyncMock() mock.db.litellm_auditlog.count = AsyncMock() + mock.replica_db = mock.db yield mock diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py index 32bcee7cb2a..d8e82280710 100644 --- a/tests/litellm_utils_tests/test_proxy_budget_reset.py +++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py @@ -141,6 +141,7 @@ async def test_reset_budget_keys_partial_failure(): key6 = {"id": "key6", "spend": 35.0, "budget_duration": 60} # Should be updated prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock( return_value=[key1, key2, key3, key4, key5, key6] ) @@ -238,6 +239,7 @@ async def test_reset_budget_users_partial_failure(): user6 = {"id": "user6", "spend": 45.0, "budget_duration": 120} # Should be updated prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock( return_value=[user1, user2, user3, user4, user5, user6] ) @@ -326,6 +328,7 @@ async def test_reset_budget_endusers_cascade_failure_is_all_or_nothing(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db async def get_data_mock(table_name, *args, **kwargs): if table_name == "budget": @@ -385,6 +388,7 @@ async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance() ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db async def get_data_mock(table_name, *args, **kwargs): if table_name == "budget": @@ -437,6 +441,7 @@ async def test_reset_budget_teams_partial_failure(): team2 = {"id": "team2", "spend": 35.0, "budget_duration": 180} # Should be updated prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(return_value=[team1, team2]) prisma_client.update_data = AsyncMock() batch_calls = _wire_batcher_for_test(prisma_client) @@ -525,6 +530,7 @@ async def test_reset_budget_continues_other_categories_on_failure(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db async def fake_get_data(*, table_name, query_type, **kwargs): if table_name == "key": @@ -663,6 +669,7 @@ async def test_service_logger_keys_success(): ), ] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(return_value=keys) prisma_client.update_data = AsyncMock() _wire_batcher_for_test(prisma_client) @@ -720,6 +727,7 @@ async def test_service_logger_keys_failure(): {"id": "key2", "spend": 15.0, "budget_duration": 60}, ] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(return_value=keys) prisma_client.update_data = AsyncMock() @@ -786,6 +794,7 @@ async def test_service_logger_users_success(): ), ] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(return_value=users) prisma_client.update_data = AsyncMock() _wire_batcher_for_test(prisma_client) @@ -839,6 +848,7 @@ async def test_service_logger_users_failure(): {"id": "user2", "spend": 25.0, "budget_duration": 120}, ] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(return_value=users) prisma_client.update_data = AsyncMock() @@ -902,6 +912,7 @@ async def test_service_logger_teams_success(): ), ] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(return_value=teams) prisma_client.update_data = AsyncMock() _wire_batcher_for_test(prisma_client) @@ -955,6 +966,7 @@ async def test_service_logger_teams_failure(): {"id": "team2", "spend": 35.0, "budget_duration": 180}, ] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(return_value=teams) prisma_client.update_data = AsyncMock() @@ -1032,6 +1044,7 @@ async def test_service_logger_endusers_success(): return [] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(side_effect=fake_get_data) prisma_client.update_data = AsyncMock() batch_calls = _wire_batcher_for_test(prisma_client) @@ -1097,6 +1110,7 @@ async def test_service_logger_endusers_failure(): return [] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.get_data = AsyncMock(side_effect=fake_get_data) prisma_client.update_data = AsyncMock() _wire_batcher_for_test(prisma_client, fail_commit=True) @@ -1154,6 +1168,7 @@ async def test_reset_budget_for_litellm_team_members_called(): enduser1 = _attrify({"user_id": "user1", "spend": 25.0, "budget_id": "budget1"}) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db async def fake_get_data(*, table_name, query_type, **kwargs): if table_name == "budget": diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py index 0a3e1a0e982..02005eeef33 100644 --- a/tests/logging_callback_tests/test_alerting.py +++ b/tests/logging_callback_tests/test_alerting.py @@ -905,6 +905,7 @@ async def test_spend_report_cache(report_type): mock_prisma.db.query_raw = AsyncMock( side_effect=[mock_spend_data, mock_tag_data] ) + mock_prisma.replica_db = mock_prisma.db slack_alerting = SlackAlerting( alerting=["webhook"], internal_usage_cache=DualCache() diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py index b6209853d82..548df6bc325 100644 --- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py +++ b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py @@ -338,6 +338,7 @@ async def test_batch_status_sync_from_provider_to_database(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( return_value=mock_db_batch ) @@ -449,6 +450,7 @@ async def test_batch_cancel_updates_database(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( return_value=None ) @@ -528,6 +530,7 @@ async def test_batch_terminal_state_skip_provider_call(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( return_value=mock_db_batch ) @@ -593,6 +596,7 @@ async def test_batch_no_status_change_skip_update(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() # Mock managed_files_obj diff --git a/tests/pass_through_unit_tests/test_claude_code_marketplace.py b/tests/pass_through_unit_tests/test_claude_code_marketplace.py index bedb8830559..bcab3d38855 100644 --- a/tests/pass_through_unit_tests/test_claude_code_marketplace.py +++ b/tests/pass_through_unit_tests/test_claude_code_marketplace.py @@ -57,6 +57,7 @@ def mock_prisma_client(): # Mock the db attribute mock_client.db = MagicMock() + mock_client.replica_db = mock_client.db # Mock the plugin table with async methods mock_table = MagicMock() diff --git a/tests/pass_through_unit_tests/test_passthrough_managed_ids.py b/tests/pass_through_unit_tests/test_passthrough_managed_ids.py index 0fc0e0e751c..8d749b2cc29 100644 --- a/tests/pass_through_unit_tests/test_passthrough_managed_ids.py +++ b/tests/pass_through_unit_tests/test_passthrough_managed_ids.py @@ -66,6 +66,7 @@ def _prisma_client() -> MagicMock: """Return a MagicMock prisma_client with async db methods.""" 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(return_value=[]) @@ -1904,6 +1905,7 @@ class TestListPassthroughIdsFromDb: pc = MagicMock() pc.db = _DbWithoutManagedTables() + pc.replica_db = pc.db for route in ("/openai/v1/files", "/openai/v1/batches"): result = await list_passthrough_ids_from_db( diff --git a/tests/unit/llms/litellm_proxy/test_skills_ownership.py b/tests/unit/llms/litellm_proxy/test_skills_ownership.py index 6caa2da3169..eaa774bd6a8 100644 --- a/tests/unit/llms/litellm_proxy/test_skills_ownership.py +++ b/tests/unit/llms/litellm_proxy/test_skills_ownership.py @@ -207,7 +207,8 @@ async def test_should_forward_skill_auth_through_transformation_handler(monkeypa async def test_should_store_team_owner_for_keys_without_user_id(monkeypatch): table = AsyncMock() table.create.side_effect = lambda data: _skill(data["skill_id"], data["created_by"]) - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( LiteLLMSkillsHandler, "_get_prisma_client", @@ -229,7 +230,8 @@ async def test_should_store_team_owner_for_keys_without_user_id(monkeypatch): async def test_should_store_token_owner_for_keys_without_user_team_or_org(monkeypatch): table = AsyncMock() table.create.side_effect = lambda data: _skill(data["skill_id"], data["created_by"]) - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( LiteLLMSkillsHandler, "_get_prisma_client", @@ -253,7 +255,8 @@ async def test_should_reject_skill_create_for_identityless_proxy_auth(monkeypatc sentinel as ``created_by`` would let any two such callers see each other's skills via the resulting shared owner scope.""" table = AsyncMock() - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( LiteLLMSkillsHandler, "_get_prisma_client", @@ -274,7 +277,8 @@ async def test_should_reject_skill_create_for_identityless_proxy_auth(monkeypatc async def test_should_filter_list_skills_to_authenticated_owner_scopes(monkeypatch): table = AsyncMock() table.find_many.return_value = [_skill("litellm_skill_owner", "user-1")] - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( LiteLLMSkillsHandler, "_get_prisma_client", @@ -299,7 +303,8 @@ async def test_should_filter_list_skills_to_authenticated_owner_scopes(monkeypat async def test_should_hide_skill_from_different_owner(monkeypatch): table = AsyncMock() table.find_unique.return_value = _skill("litellm_skill_other", "user-2") - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( LiteLLMSkillsHandler, "_get_prisma_client", @@ -319,7 +324,8 @@ async def test_should_hide_skill_from_different_owner(monkeypatch): async def test_should_hide_unowned_skill_by_default(monkeypatch): table = AsyncMock() table.find_unique.return_value = _skill("litellm_skill_unowned", None) - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( LiteLLMSkillsHandler, "_get_prisma_client", @@ -341,7 +347,8 @@ async def test_list_skills_excludes_unowned_for_non_admin(monkeypatch): with ``created_by IS NULL`` are excluded — admin-only.""" table = AsyncMock() table.find_many.return_value = [] - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( LiteLLMSkillsHandler, "_get_prisma_client", @@ -397,7 +404,8 @@ async def test_load_skill_uses_cache_after_first_db_hit(monkeypatch): fake_skill = Mock(created_by="user-1", skill_id="litellm_skill_a") table = AsyncMock() table.find_unique = AsyncMock(return_value=fake_skill) - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( skills_handler.LiteLLMSkillsHandler, "_get_prisma_client", @@ -415,7 +423,8 @@ async def test_load_skill_caches_negative_lookups(monkeypatch): the DB and the caller still sees ``None``.""" table = AsyncMock() table.find_unique = AsyncMock(return_value=None) - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( skills_handler.LiteLLMSkillsHandler, "_get_prisma_client", @@ -434,7 +443,8 @@ async def test_delete_skill_invalidates_cache(monkeypatch): table = AsyncMock() table.find_unique = AsyncMock(return_value=fake_skill) table.delete = AsyncMock() - prisma_client = type("Prisma", (), {"db": type("DB", (), {"litellm_skillstable": table})()})() + db = type("DB", (), {"litellm_skillstable": table})() + prisma_client = type("Prisma", (), {"db": db, "replica_db": db})() monkeypatch.setattr( skills_handler.LiteLLMSkillsHandler, "_get_prisma_client", diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py index e185d95ffb8..64c7e29bf17 100644 --- a/tests/unit/repositories/test_repositories.py +++ b/tests/unit/repositories/test_repositories.py @@ -136,6 +136,7 @@ class MockPrismaClient: self.db.litellm_projecttable = MockTable(pk_field="project_id") self.db.litellm_objectpermissiontable = MockTable(pk_field="object_permission_id") self.db.litellm_credentialstable = MockTable() + self.replica_db = self.db class TestBaseRepository: @@ -303,9 +304,8 @@ class TestModelRepository: @pytest.mark.asyncio async def test_find_all_except_serializes_exclusion_for_prisma(self) -> None: find_many: Final = AsyncMock(return_value=[]) - client: Final = SimpleNamespace( - db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many)) - ) + db: Final = SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many)) + client: Final = SimpleNamespace(db=db, replica_db=db) await ModelRepository(client).find_all_except("current-model") @@ -2090,6 +2090,7 @@ class TestPrismaTableRepository: ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db agents = AgentsRepository(prisma_client) policy = PolicyRepository(prisma_client) @@ -2131,6 +2132,7 @@ class TestPrismaTableRepository: ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db repos = [ obj for name, obj in vars(tr).items() @@ -2263,6 +2265,7 @@ class TestAutoRouterSessionRepository: client = MagicMock() client.db.litellm_autoroutersession = _Table() + client.replica_db = client.db return AutoRouterSessionRepository(client), lookups @pytest.mark.asyncio @@ -2287,6 +2290,7 @@ class TestAutoRouterSessionRepository: from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository client = MagicMock() + client.replica_db = client.db assert AutoRouterSessionRepository(client).table is client.db.litellm_autoroutersession with pytest.raises(RuntimeError, match="No DB Connected"): _ = AutoRouterSessionRepository(None).table