From f6f392033708254f20ee7ba755a405096c6a07bc Mon Sep 17 00:00:00 2001 From: yuneng Date: Thu, 24 Sep 2026 08:14:48 +0000 Subject: [PATCH] refactor(db): make read-replica routing an explicit opt-in via PrismaClient.replica_db Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../enterprise_hooks/blocked_user_list.py | 2 +- .../send_emails/base_email.py | 4 +- .../send_emails/endpoints.py | 4 +- .../proxy/common_utils/check_batch_cost.py | 8 +- .../common_utils/check_responses_cost.py | 2 +- .../proxy/hooks/managed_files.py | 4 +- .../proxy/vector_stores/endpoints.py | 2 +- .../SlackAlerting/user_spend_alerts.py | 4 +- litellm/integrations/cloudzero/database.py | 2 +- litellm/integrations/email_alerting.py | 2 +- litellm/integrations/focus/database.py | 4 +- litellm/integrations/shadow_eval_logger.py | 4 +- .../base_managed_resource.py | 2 +- .../mcp_server/oauth2_flow_backfill.py | 4 +- .../mcp_server/oauth_issuer_stamp_backfill.py | 2 +- .../sso_assertion_store.py | 4 +- litellm/proxy/agent_endpoints/endpoints.py | 2 +- .../analytics_endpoints/cache_activity.py | 12 +- litellm/proxy/auth/auth_object_prefetch.py | 2 +- litellm/proxy/auth/user_api_key_auth.py | 2 +- .../proxy/common_utils/reset_budget_job.py | 2 +- .../proxy/db/budget_window_spend_writer.py | 6 +- litellm/proxy/db/health_check_latest.py | 4 +- litellm/proxy/db/routing_prisma_wrapper.py | 4 + .../auto_router_endpoints.py | 14 +- .../common_daily_activity.py | 8 +- .../credential_migration.py | 14 +- .../gateway_request_endpoints.py | 2 +- .../internal_user_endpoints.py | 2 +- .../key_management_endpoints.py | 8 +- .../management_v1/budgets.py | 4 +- .../management_v1/spend_logs.py | 2 +- .../model_management_endpoints.py | 6 +- .../management_endpoints/router_weights.py | 4 +- .../management_endpoints/team_endpoints.py | 2 +- .../user_agent_analytics_endpoints.py | 2 +- litellm/proxy/proxy_server.py | 10 +- .../daily_global_spend_rollup.py | 6 +- .../spend_tracking/key_metadata_recovery.py | 2 +- .../spend_management_endpoints.py | 4 +- .../spend_tracking/spend_tracking_utils.py | 2 +- litellm/proxy/utils.py | 29 ++- .../autorouter_session_repository.py | 2 +- litellm/repositories/budget_repository.py | 6 +- litellm/repositories/config_repository.py | 2 +- .../repositories/credentials_repository.py | 4 +- .../object_permission_repository.py | 2 +- .../repositories/organization_repository.py | 6 +- litellm/repositories/prisma_protocols.py | 4 +- litellm/repositories/project_repository.py | 2 +- litellm/repositories/table_repositories.py | 2 +- litellm/repositories/team_repository.py | 4 +- litellm/repositories/user_repository.py | 2 +- .../verification_token_repository.py | 4 +- .../session_handler.py | 2 +- .../SlackAlerting/test_user_spend_alerts.py | 2 + .../integrations/test_shadow_eval_logger.py | 1 + .../test_sso_assertion_store.py | 3 + .../mcp_server/test_oauth2_flow_backfill.py | 1 + .../test_oauth_issuer_stamp_backfill.py | 1 + .../proxy/agent_endpoints/test_endpoints.py | 11 + .../proxy/auth/test_auth_object_prefetch.py | 1 + .../proxy/auth/test_user_api_key_auth.py | 20 +- .../common_utils/test_reset_budget_job.py | 9 + .../db/test_budget_window_spend_writer.py | 1 + .../proxy/db/test_health_check_latest.py | 1 + .../management_v1/test_budgets.py | 1 + .../management_v1/test_spend_logs.py | 1 + .../test_auto_router_endpoints.py | 22 +- .../test_common_daily_activity.py | 28 +++ .../test_credential_migration.py | 13 + .../test_gateway_request_endpoints.py | 1 + .../test_internal_user_endpoints.py | 66 +++++ .../test_key_management_endpoints.py | 234 +++++++++++++++++- .../test_model_management_endpoints.py | 48 +++- .../test_team_endpoints.py | 168 ++++++++++++- .../test_daily_global_spend_rollup.py | 1 + .../test_key_metadata_recovery.py | 28 +++ .../test_spend_management_endpoints.py | 50 +++- tests/test_litellm/proxy/test_proxy_server.py | 113 +++++++++ tests/test_litellm/proxy/test_proxy_utils.py | 104 ++++++++ .../test_session_handler.py | 1 + 82 files changed, 1050 insertions(+), 131 deletions(-) diff --git a/enterprise/enterprise_hooks/blocked_user_list.py b/enterprise/enterprise_hooks/blocked_user_list.py index a032ea7662d..39efab054a1 100644 --- a/enterprise/enterprise_hooks/blocked_user_list.py +++ b/enterprise/enterprise_hooks/blocked_user_list.py @@ -90,7 +90,7 @@ class _ENTERPRISE_BlockedUserList(CustomLogger): if end_user_cache_obj is None and self.prisma_client is not None: # check db end_user_obj = ( - await self.prisma_client.db.litellm_endusertable.find_unique( + await self.prisma_client.replica_db.litellm_endusertable.find_unique( where={"user_id": user} ) ) diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py index 4be09670e92..439462d653a 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -813,7 +813,7 @@ class BaseEmailLogger(CustomLogger): ) return None - user_row = await prisma_client.db.litellm_usertable.find_unique( + user_row = await prisma_client.replica_db.litellm_usertable.find_unique( where={"user_id": user_id} ) @@ -898,7 +898,7 @@ class BaseEmailLogger(CustomLogger): try: # Try to get existing invitation existing_invitations = ( - await prisma_client.db.litellm_invitationlink.find_many( + await prisma_client.replica_db.litellm_invitationlink.find_many( where={"user_id": user_id}, order={"created_at": "desc"}, ) diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py index 1ab173a915a..c81f234daa5 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/endpoints.py @@ -25,7 +25,7 @@ async def _get_email_settings(prisma_client) -> Dict[str, bool]: """Helper function to get email settings from general_settings in db""" try: # Get general settings from db - general_settings_entry = await prisma_client.db.litellm_config.find_unique( + general_settings_entry = await prisma_client.replica_db.litellm_config.find_unique( where={"param_name": "general_settings"} ) @@ -71,7 +71,7 @@ async def _save_email_settings(prisma_client, settings: Dict[str, bool]): ) # Get current general settings - general_settings_entry = await prisma_client.db.litellm_config.find_unique( + general_settings_entry = await prisma_client.replica_db.litellm_config.find_unique( where={"param_name": "general_settings"} ) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 41974c26158..0963770dc28 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -57,24 +57,24 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable + table: Final[TableActions[_ManagedObjectRow]] = prisma_client.replica_db.litellm_managedobjecttable return table def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]": - table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable + table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.replica_db.litellm_usertable return table def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]": table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = ( - prisma_client.db.litellm_verificationtoken + prisma_client.replica_db.litellm_verificationtoken ) return table def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]": - table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable + table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.replica_db.litellm_teamtable return table diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index cdeea0d3d4b..65a49ba089a 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -43,7 +43,7 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable + table: Final[TableActions[_ManagedObjectRow]] = prisma_client.replica_db.litellm_managedobjecttable return table diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 5ac7c1e53c1..afce5dec6ec 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -197,11 +197,11 @@ class _CursorPageArgs(TypedDict, total=False): def _managed_file_table(prisma_client: PrismaClient) -> _ManagedFileTableActions: - return prisma_client.db.litellm_managedfiletable + return prisma_client.replica_db.litellm_managedfiletable def _managed_object_table(prisma_client: PrismaClient) -> _ManagedObjectTableActions: - return prisma_client.db.litellm_managedobjecttable + return prisma_client.replica_db.litellm_managedobjecttable def _storage_metadata_of(file_object: OpenAIFileObject | None) -> Mapping[str, str]: diff --git a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py index e95a7c99971..f627c0b26ec 100644 --- a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py +++ b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py @@ -59,7 +59,7 @@ class ManagedVectorStoreTable(Protocol): def managed_vector_store_table(prisma_client: "PrismaClient") -> ManagedVectorStoreTable: """The Prisma table actions for managed vector stores, behind a typed surface.""" - return prisma_client.db.litellm_managedvectorstorestable + return prisma_client.replica_db.litellm_managedvectorstorestable ######################################################## diff --git a/litellm/integrations/SlackAlerting/user_spend_alerts.py b/litellm/integrations/SlackAlerting/user_spend_alerts.py index 38794735c1b..c98948b1785 100644 --- a/litellm/integrations/SlackAlerting/user_spend_alerts.py +++ b/litellm/integrations/SlackAlerting/user_spend_alerts.py @@ -56,7 +56,9 @@ async def fetch_user_spend_rows( today_str: Final = today.strftime("%Y-%m-%d") month_start_str: Final = today.replace(day=1).strftime("%Y-%m-%d") baseline_start_str: Final = (today - datetime.timedelta(days=max(baseline_days, 1))).strftime("%Y-%m-%d") - raw: Final = await prisma_client.db.query_raw(USER_SPEND_QUERY, today_str, month_start_str, baseline_start_str) + raw: Final = await prisma_client.replica_db.query_raw( + USER_SPEND_QUERY, today_str, month_start_str, baseline_start_str + ) return USER_SPEND_ROWS_ADAPTER.validate_python(raw) diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index 2fb10ad8a96..17128187bb3 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -93,7 +93,7 @@ class LiteLLMDatabase: query += " LIMIT $3" try: - db_response: Final = await client.db.query_raw(query, *params) + db_response: Final = await client.replica_db.query_raw(query, *params) from litellm.proxy.spend_tracking.key_metadata_recovery import ( fill_missing_api_key_aliases, ) diff --git a/litellm/integrations/email_alerting.py b/litellm/integrations/email_alerting.py index 351896425bb..771afb387ef 100644 --- a/litellm/integrations/email_alerting.py +++ b/litellm/integrations/email_alerting.py @@ -51,7 +51,7 @@ async def get_all_team_member_emails(team_id: str | None = None) -> list: WHERE user_id = ANY($1::TEXT[]); """ - _result: Final = await prisma_client.db.query_raw(sql_query, _team_member_user_ids) + _result: Final = await prisma_client.replica_db.query_raw(sql_query, _team_member_user_ids) verbose_logger.debug("Email Alerting: Got all Emails for team, emails=%s", _result) diff --git a/litellm/integrations/focus/database.py b/litellm/integrations/focus/database.py index 891318f1c54..3f60a1e9265 100644 --- a/litellm/integrations/focus/database.py +++ b/litellm/integrations/focus/database.py @@ -95,7 +95,7 @@ class FocusLiteLLMDatabase: """ try: - db_response: Final = await client.db.query_raw(query, *query_params) + db_response: Final = await client.replica_db.query_raw(query, *query_params) from litellm.proxy.spend_tracking.key_metadata_recovery import ( fill_missing_api_key_aliases, ) @@ -123,7 +123,7 @@ class FocusLiteLLMDatabase: ORDER BY ordinal_position; """ try: - columns_response: Final = await client.db.query_raw(info_query) + columns_response: Final = await client.replica_db.query_raw(info_query) return {"columns": columns_response, "table_name": "LiteLLM_DailyUserSpend"} except Exception as exc: raise RuntimeError(f"Error getting table info: {exc}") from exc diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index 19d9bee7493..3436285ee36 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -863,14 +863,14 @@ class ShadowEvalLogger(CustomLogger): if prisma is None: return _EMPTY_JOBS try: - records: Final = await prisma.db.litellm_shadowevaljob.find_many( + records: Final = await prisma.replica_db.litellm_shadowevaljob.find_many( where={ # mutable-ok: Prisma filter "stopped_at": None, "ends_at": {"gt": datetime.now(timezone.utc)}, # mutable-ok: Prisma filter }, ) grouped: Final = ( - await prisma.db.litellm_shadowevalattempt.group_by( + await prisma.replica_db.litellm_shadowevalattempt.group_by( by=["job_id"], count=True, sum={"judge_cost": True, "shadow_cost": True, "shadow_classifier_cost": True}, diff --git a/litellm/llms/base_llm/managed_resources/base_managed_resource.py b/litellm/llms/base_llm/managed_resources/base_managed_resource.py index 4fbc0ce51b0..f3af18e12ac 100644 --- a/litellm/llms/base_llm/managed_resources/base_managed_resource.py +++ b/litellm/llms/base_llm/managed_resources/base_managed_resource.py @@ -90,7 +90,7 @@ class BaseManagedResource(ABC, Generic[ResourceObjectType]): self.prisma_client = prisma_client def _resource_table(self) -> _ManagedResourceTable[ResourceObjectType]: - return getattr(self.prisma_client.db, self.table_name) + return getattr(self.prisma_client.replica_db, self.table_name) # ============================================================================ # ABSTRACT METHODS diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py index 150900e7ff2..d0c01f3efbb 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py @@ -99,12 +99,12 @@ class _MCPUserCredentialsTable(Protocol): def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable: """The MCP server table, typed so the untyped prisma client surface stops here.""" - return prisma_client.db.litellm_mcpservertable + return prisma_client.replica_db.litellm_mcpservertable def _mcp_user_credentials_table(prisma_client: PrismaClient) -> _MCPUserCredentialsTable: """The per-user MCP credential table, typed so the untyped prisma client surface stops here.""" - return prisma_client.db.litellm_mcpusercredentials + return prisma_client.replica_db.litellm_mcpusercredentials def _decrypted_credentials(raw_credentials: str | Mapping[str, JsonValue] | None) -> MCPCredentials | None: diff --git a/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py index 62f4489a784..db56455a7ee 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py @@ -105,7 +105,7 @@ def _is_stamped_issuer_row(row: _MCPServerRow) -> bool: async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int: """Clear gateway-written issuer stamps, returning the number of rows healed.""" - candidate_rows: Final[list[_MCPServerRow]] = await prisma_client.db.litellm_mcpservertable.find_many( + candidate_rows: Final[list[_MCPServerRow]] = await prisma_client.replica_db.litellm_mcpservertable.find_many( where={ "updated_by": _DISCOVERY_ACTOR, "auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)}, diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py index 5503d19211b..5cd49fb6b9f 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py @@ -59,12 +59,12 @@ class _MCPServerTable(Protocol): def _assertion_table(prisma_client: PrismaClient) -> _SSOAssertionTable: """The SSO assertion table, typed so the untyped prisma client surface stops here.""" - return prisma_client.db.litellm_ssoidentityassertion + return prisma_client.replica_db.litellm_ssoidentityassertion def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable: """The MCP server table, typed so the untyped prisma client surface stops here.""" - return prisma_client.db.litellm_mcpservertable + return prisma_client.replica_db.litellm_mcpservertable class SSOIdentityAssertion(BaseModel): diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index aa8979a73c6..41ad7790ef2 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -121,7 +121,7 @@ async def _attach_keys_to_agents(agents: Sequence[AgentResponse], prisma_client) ) if not agent_ids: return - key_rows: Final = await prisma_client.db.litellm_verificationtoken.find_many( + key_rows: Final = await prisma_client.replica_db.litellm_verificationtoken.find_many( where={"agent_id": {"in": agent_ids}}, ) keys_by_agent: Final[dict[str, list[AgentKeySummary]]] = {} diff --git a/litellm/proxy/analytics_endpoints/cache_activity.py b/litellm/proxy/analytics_endpoints/cache_activity.py index 5de3b610782..83119d0adfc 100644 --- a/litellm/proxy/analytics_endpoints/cache_activity.py +++ b/litellm/proxy/analytics_endpoints/cache_activity.py @@ -22,7 +22,7 @@ class _SupportsRawQueryDb(Protocol): """A prisma client handle, narrowed to the raw-query surface used here.""" @property - def db(self) -> _SupportsQueryRaw: ... + def replica_db(self) -> _SupportsQueryRaw: ... class CacheActivityGroup(BaseModel): @@ -169,12 +169,14 @@ async def get_cache_activity( key_aliases_json: Final = json.dumps(list(key_aliases)) models_json: Final = json.dumps(list(models)) group_rows, error_rows, key_alias_rows, model_rows = await asyncio.gather( - prisma_client.db.query_raw(GROUPS_SQL, start_date, end_date, key_aliases_json, models_json, INFO_ROUTES_JSON), - prisma_client.db.query_raw( + prisma_client.replica_db.query_raw( + GROUPS_SQL, start_date, end_date, key_aliases_json, models_json, INFO_ROUTES_JSON + ), + prisma_client.replica_db.query_raw( ERROR_BREAKDOWN_SQL, start_date, end_date, key_aliases_json, models_json, INFO_ROUTES_JSON ), - prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON), - prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON), + prisma_client.replica_db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON), + prisma_client.replica_db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date, INFO_ROUTES_JSON), ) groups: Final = _groups_adapter.validate_python(group_rows or []) return CacheActivityResponse( diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index 52e26e885c9..7647e7b562e 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -248,7 +248,7 @@ def _validate_row( async def _fetch_rows( refs: AuthObjectRefs, kinds: frozenset[_RowKind], prisma_client: PrismaClient ) -> Mapping[str, object]: - row: Final[object] = await prisma_client.db.query_first( # pyright: ignore[reportAny] # prisma types query_first as Any + row: Final[object] = await prisma_client.replica_db.query_first( # pyright: ignore[reportAny] # prisma types query_first as Any _SQL, refs.user_id if "user_row" in kinds else None, refs.team_id if kinds & _TEAM_BOUND_ROWS else None, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a6c0792a86f..d4e3b224273 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -735,7 +735,7 @@ async def _fetch_global_spend_with_event_coordination( """ async def _load_global_spend() -> float | None: - proxy_budget_row: Final = await prisma_client.db.litellm_usertable.find_unique( + proxy_budget_row: Final = await prisma_client.replica_db.litellm_usertable.find_unique( where={"user_id": LITELLM_PROXY_BUDGET_NAME} ) return float(proxy_budget_row.spend) if proxy_budget_row is not None else None diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index b35b876b475..3ece35bf210 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -1524,7 +1524,7 @@ class ResetBudgetJob: ) -> str | None: """Reset one page of windows; return the next cursor, or None when drained.""" rows: Final = await self._with_db_retry( - lambda: self.prisma_client.db.query_raw(source.page_query(), cursor, RESET_BUDGET_JOB_BATCH_SIZE), + lambda: self.prisma_client.replica_db.query_raw(source.page_query(), cursor, RESET_BUDGET_JOB_BATCH_SIZE), reason=f"reset_budget_read_{source.retry_subject}_windows_failure", ) for row in rows: diff --git a/litellm/proxy/db/budget_window_spend_writer.py b/litellm/proxy/db/budget_window_spend_writer.py index 8cf2f737063..12b16a7acf9 100644 --- a/litellm/proxy/db/budget_window_spend_writer.py +++ b/litellm/proxy/db/budget_window_spend_writer.py @@ -147,9 +147,9 @@ async def spend_logs_seed_totals( else: return None rows: Final = ( - await prisma_client.db.query_raw(unbounded_sql, entity_id, window_start) + await prisma_client.replica_db.query_raw(unbounded_sql, entity_id, window_start) if batch_started_at is None - else await prisma_client.db.query_raw( + else await prisma_client.replica_db.query_raw( bounded_sql, entity_id, window_start, @@ -182,7 +182,7 @@ async def _existing_primary_keys( prisma_client: "PrismaClient", transactions: tuple[WindowSpendTransaction, ...], ) -> frozenset[tuple[str, str, str]]: - rows: Final = await prisma_client.db.query_raw( + rows: Final = await prisma_client.replica_db.query_raw( _SELECT_EXISTING_ROWS_SQL, tuple(transaction["entity_type"] for transaction in transactions), tuple(transaction["entity_id"] for transaction in transactions), diff --git a/litellm/proxy/db/health_check_latest.py b/litellm/proxy/db/health_check_latest.py index 21438f095bb..ad352c7de60 100644 --- a/litellm/proxy/db/health_check_latest.py +++ b/litellm/proxy/db/health_check_latest.py @@ -75,7 +75,7 @@ _ROWS_ADAPTER: Final = TypeAdapter(tuple[LatestHealthCheckRow, ...]) async def query_latest_health_checks(prisma_client: PrismaClient) -> tuple[LatestHealthCheckRow, ...]: - rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL) + rows: Final = await prisma_client.replica_db.query_raw(LATEST_HEALTH_CHECKS_SQL) return _ROWS_ADAPTER.validate_python(rows) @@ -93,7 +93,7 @@ async def fetch_latest_health_checks_for_models( if not model_names: return () try: - rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, list(model_names)) + rows: Final = await prisma_client.replica_db.query_raw(LATEST_HEALTH_CHECKS_FOR_MODELS_SQL, list(model_names)) return _ROWS_ADAPTER.validate_python(rows) except Exception as query_err: # noqa: BLE001 # a paged model list must not fail on its health decoration verbose_proxy_logger.error("Error getting latest health checks for models: %s", query_err) diff --git a/litellm/proxy/db/routing_prisma_wrapper.py b/litellm/proxy/db/routing_prisma_wrapper.py index 0eb378b2fe9..de60ac81217 100644 --- a/litellm/proxy/db/routing_prisma_wrapper.py +++ b/litellm/proxy/db/routing_prisma_wrapper.py @@ -80,6 +80,10 @@ class WriterPinnedClient: def __init__(self, db: "PrismaWrapper | RoutingPrismaWrapper") -> None: self.db: Final = db.writer if isinstance(db, RoutingPrismaWrapper) and not db.writer_unavailable else db + @property + def replica_db(self) -> "PrismaWrapper | RoutingPrismaWrapper": + return self.db + def writer_wrapper(db: "PrismaWrapper | RoutingPrismaWrapper") -> PrismaWrapper: """Unlike `WriterPinnedClient`, ignores `writer_unavailable`: a raw SQL write has no replica fallback.""" diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index 9708161397a..fd445823a3f 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -183,31 +183,31 @@ def _team_table(prisma_client: "PrismaClient") -> _TeamTable: def _verification_tokens(prisma_client: "PrismaClient") -> _VerificationTokenTable: - return prisma_client.db.litellm_verificationtoken + return prisma_client.replica_db.litellm_verificationtoken def _team_rows(prisma_client: "PrismaClient") -> _TeamRowsTable: - return prisma_client.db.litellm_teamtable + return prisma_client.replica_db.litellm_teamtable def _user_rows(prisma_client: "PrismaClient") -> _UserRowsTable: - return prisma_client.db.litellm_usertable + return prisma_client.replica_db.litellm_usertable def _shadow_eval_jobs(prisma_client: "PrismaClient") -> _ShadowEvalJobTable: - return prisma_client.db.litellm_shadowevaljob + return prisma_client.replica_db.litellm_shadowevaljob def _shadow_eval_funnel(prisma_client: "PrismaClient") -> _ShadowEvalFunnelTable: - return prisma_client.db.litellm_shadowevalfunnel # pyright: ignore[reportAttributeAccessIssue] # generated client + return prisma_client.replica_db.litellm_shadowevalfunnel # pyright: ignore[reportAttributeAccessIssue] # generated client def _shadow_eval_attempts(prisma_client: "PrismaClient") -> _ShadowEvalAttemptTable: - return prisma_client.db.litellm_shadowevalattempt + return prisma_client.replica_db.litellm_shadowevalattempt async def _query_raw(prisma_client: "PrismaClient", query: str, *args: object) -> Sequence[Mapping[str, object]]: - return await prisma_client.db.query_raw(query, *args) + return await prisma_client.replica_db.query_raw(query, *args) async def _authorize_router_dry_run(user_api_key_dict: UserAPIKeyAuth, team_id: str | None) -> LiteLLM_TeamTable | None: diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index bf8e7bc15fb..a973a89915f 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -206,7 +206,7 @@ async def _query_raw_optional( ) -> list[dict[str, object]] | None: # mutable-ok: prisma query_raw return shape if query is None: return None - return await prisma_client.db.query_raw(query[0], *query[1]) + return await prisma_client.replica_db.query_raw(query[0], *query[1]) def _reported_flat_cost(record: DailySpendRecord | _RollupMetricsRow) -> float: @@ -832,7 +832,7 @@ def _build_aggregated_sql_query( """Build the GROUPING SETS query for aggregated daily activity. Returns: - Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw(). + Tuple of (sql_query, params_list) ready for prisma_client.replica_db.query_raw(). """ pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name) if pg_table is None: @@ -1306,7 +1306,7 @@ async def get_daily_activity( include_current_utc_day=include_current_utc_day, ) - spend_table: Final[TableActions[DailySpendRecord]] = getattr(prisma_client.db, table_name) + spend_table: Final[TableActions[DailySpendRecord]] = getattr(prisma_client.replica_db, table_name) # Get total count for pagination total_count: Final[int] = await spend_table.count(where=where_conditions) @@ -1471,7 +1471,7 @@ async def get_daily_activity_aggregated( entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None raw_rows, raw_entity_rows = await asyncio.gather( - prisma_client.db.query_raw(sql_query, *sql_params), + prisma_client.replica_db.query_raw(sql_query, *sql_params), _query_raw_optional(prisma_client, entity_query), ) diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index 915cce87dbd..4763cf4986d 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -222,7 +222,7 @@ async def _migrate_config_settings_row( dict with selected sensitive fields (vantage_settings / cloudzero_settings). """ report: Final = LocationReport(location=param_name) - record: Final = await prisma_client.db.litellm_config.find_unique(where={"param_name": param_name}) + record: Final = await prisma_client.replica_db.litellm_config.find_unique(where={"param_name": param_name}) if record is None or record.param_value is None: return report @@ -274,7 +274,7 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR every present string field. """ report: Final = LocationReport(location="sso_config") - record: Final = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"}) + record: Final = await prisma_client.replica_db.litellm_ssoconfig.find_unique(where={"id": "sso_config"}) if record is None or record.sso_settings is None: return report @@ -341,10 +341,10 @@ async def _migrate_callback_vars_table( report: Final = LocationReport(location=f"{table_name}.callback_vars") if table_name == "team": - table = prisma_client.db.litellm_teamtable + table = prisma_client.replica_db.litellm_teamtable pk = "team_id" else: - table = prisma_client.db.litellm_verificationtoken + table = prisma_client.replica_db.litellm_verificationtoken pk = "token" rows: Final = await table.find_many() @@ -510,7 +510,7 @@ async def _scan_one_table( scalar_columns: tuple, ) -> LocationReport: report: Final = LocationReport(location=location) - table: Final = getattr(prisma_client.db, db_attr, None) + table: Final = getattr(prisma_client.replica_db, db_attr, None) if table is None: return report try: @@ -544,7 +544,9 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport: """Scan the ``environment_variables`` config row (``param_value`` dict).""" report: Final = LocationReport(location="config_environment_variables") try: - record: Final = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"}) + record: Final = await prisma_client.replica_db.litellm_config.find_unique( + where={"param_name": "environment_variables"} + ) except Exception as e: # pragma: no cover - defensive verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e)) return report diff --git a/litellm/proxy/management_endpoints/gateway_request_endpoints.py b/litellm/proxy/management_endpoints/gateway_request_endpoints.py index 33c078274fb..a63720b43da 100644 --- a/litellm/proxy/management_endpoints/gateway_request_endpoints.py +++ b/litellm/proxy/management_endpoints/gateway_request_endpoints.py @@ -122,7 +122,7 @@ async def get_gateway_daily_activity( raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) default_start, default_end = _default_range() - raw_rows: Final = await prisma_client.db.query_raw( # pyright: ignore[reportAny] # untyped prisma client + raw_rows: Final = await prisma_client.replica_db.query_raw( # pyright: ignore[reportAny] # untyped prisma client _AGGREGATE_SQL, start_date or default_start, end_date or default_end, diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 59d8dd821d8..050b181ff2b 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -1169,7 +1169,7 @@ async def user_info_v2( async def _fetch_admin_teams_and_keys_rows( prisma_client: "PrismaClient", sql_query: str ) -> Sequence[Mapping[str, Sequence[Mapping[str, object]] | None]]: - return await prisma_client.db.query_raw(sql_query) + return await prisma_client.replica_db.query_raw(sql_query) async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth): diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 306ea90d7f1..342bea6947a 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -6733,7 +6733,9 @@ async def key_aliases( where_sql: Final = " AND ".join(where_parts) count_sql: Final = f'SELECT COUNT(*) AS count FROM "LiteLLM_VerificationToken" WHERE {where_sql}' - count_rows: Final[Sequence[Mapping[str, int]]] = await prisma_client.db.query_raw(count_sql, *query_params) + count_rows: Final[Sequence[Mapping[str, int]]] = await prisma_client.replica_db.query_raw( + count_sql, *query_params + ) total_count: Final = int(count_rows[0]["count"]) if count_rows else 0 aliases_params: Final = query_params + [size, (page - 1) * size] @@ -6746,7 +6748,9 @@ async def key_aliases( f" ORDER BY key_alias ASC" f" LIMIT ${limit_idx} OFFSET ${offset_idx}" ) - alias_rows: Final[Sequence[Mapping[str, str]]] = await prisma_client.db.query_raw(aliases_sql, *aliases_params) + alias_rows: Final[Sequence[Mapping[str, str]]] = await prisma_client.replica_db.query_raw( + aliases_sql, *aliases_params + ) aliases: Final[list[str]] = [row["key_alias"] for row in alias_rows if row.get("key_alias")] total_pages: Final = -(-total_count // size) if total_count > 0 else 0 diff --git a/litellm/proxy/management_endpoints/management_v1/budgets.py b/litellm/proxy/management_endpoints/management_v1/budgets.py index 106dbfaf7b7..b82b0ff5196 100644 --- a/litellm/proxy/management_endpoints/management_v1/budgets.py +++ b/litellm/proxy/management_endpoints/management_v1/budgets.py @@ -85,7 +85,7 @@ class PrismaBudgetListExecutor: async def count(self, where: tuple[Predicate, ...]) -> int: clauses, params = where_sql(where) sql: Final = f"SELECT COUNT(*) AS count FROM {BUDGET_TABLE}" + (f" WHERE {clauses}" if clauses else "") - rows: Final = await self.prisma_client.db.query_raw(sql, *params) + rows: Final = await self.prisma_client.replica_db.query_raw(sql, *params) counted: Final = _ROW_COUNTS.validate_python(rows) return counted[0].count if counted else 0 @@ -97,7 +97,7 @@ class PrismaBudgetListExecutor: + f" ORDER BY {order_by_sql(plan.order)}" + f" LIMIT ${len(params) + 1} OFFSET ${len(params) + 2}" ) - rows: Final = await self.prisma_client.db.query_raw(sql, *params, plan.take, plan.skip) + rows: Final = await self.prisma_client.replica_db.query_raw(sql, *params, plan.take, plan.skip) return _BUDGET_ROWS.validate_python(rows) diff --git a/litellm/proxy/management_endpoints/management_v1/spend_logs.py b/litellm/proxy/management_endpoints/management_v1/spend_logs.py index 1cbc454ca5e..9d6e37ed765 100644 --- a/litellm/proxy/management_endpoints/management_v1/spend_logs.py +++ b/litellm/proxy/management_endpoints/management_v1/spend_logs.py @@ -143,7 +143,7 @@ async def _list_spend_log_facet( f" ORDER BY {column_sql} ASC" f" LIMIT ${scan_idx + 1} OFFSET ${scan_idx + 2}" ) - rows: Final = await prisma_client.db.query_raw(facet_sql, *params) + rows: Final = await prisma_client.replica_db.query_raw(facet_sql, *params) values: Final[list[str]] = [row[column] for row in rows if row.get(column)] has_more: Final = len(values) > page_size diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ae294871afc..2f68bde22fc 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -223,6 +223,10 @@ class _ModelTransactionClient(BaseModel): class _TransactionClient: db: _TxModelTables + @property + def replica_db(self) -> _TxModelTables: + return self.db + _RowT = TypeVar("_RowT") @@ -260,7 +264,7 @@ def _repo_team_table(prisma_client: PrismaClient) -> _TeamLookupTable: def _db_team_table(prisma_client: PrismaClient) -> _TeamTable: - return prisma_client.db.litellm_teamtable + return prisma_client.replica_db.litellm_teamtable def _model_alias_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_ModelTable]": diff --git a/litellm/proxy/management_endpoints/router_weights.py b/litellm/proxy/management_endpoints/router_weights.py index 99b368808c3..aa49b4d6998 100644 --- a/litellm/proxy/management_endpoints/router_weights.py +++ b/litellm/proxy/management_endpoints/router_weights.py @@ -26,7 +26,7 @@ class _ModelDb(Protocol): class _PrismaClient(Protocol): @property @abstractmethod - def db(self) -> _ModelDb: + def replica_db(self) -> _ModelDb: pass @@ -112,7 +112,7 @@ async def validate_router_settings_weights( return if prisma_client is None: raise HTTPException(status_code=503, detail="Database unavailable while validating router weights") - stored_models: Final = await prisma_client.db.litellm_proxymodeltable.find_many( + stored_models: Final = await prisma_client.replica_db.litellm_proxymodeltable.find_many( where={"model_id": {"in": list(deployment_ids)}} ) stored_by_id: Final = { diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 6cec3e714ec..b328064eb4b 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -6894,7 +6894,7 @@ async def get_team_spend_by_user( own_user_only: Final = scope.api_key_filter is not None user_param: Final = (user_api_key_dict.user_id or "",) if own_user_only else () - rows: Final[Sequence[_TeamUserSpendDbRow]] = await prisma_client.db.query_raw( + rows: Final[Sequence[_TeamUserSpendDbRow]] = await prisma_client.replica_db.query_raw( _team_user_spend_sql(team_count=len(scoped_team_ids), restrict_to_user=own_user_only), start_date, end_date, diff --git a/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py b/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py index 0422c72cdb3..ff7a3c9c50a 100644 --- a/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py +++ b/litellm/proxy/management_endpoints/user_agent_analytics_endpoints.py @@ -156,7 +156,7 @@ def _typed_table(repo: DailyTagSpendRepository | VerificationTokenRepository | U async def _query_raw(prisma_client: "PrismaClient", sql_query: str, *params: object) -> object: - return await prisma_client.db.query_raw(sql_query, *params) + return await prisma_client.replica_db.query_raw(sql_query, *params) @router.get( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7e18db742cd..cb5c41705c2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -10124,7 +10124,7 @@ class ProxyStartupEvent: from prisma.errors import UniqueViolationError try: - config_table: Final = prisma_client.db.litellm_config + config_table: Final = prisma_client.replica_db.litellm_config row: Final = await config_table.find_unique( where={"param_name": TUNING_BASELINE_PARAM_NAME} # mutable-ok: Prisma rejects mappingproxy input ) @@ -15143,7 +15143,7 @@ async def model_streaming_metrics( """ _all_api_bases: Final = set() - db_response: Final[Sequence[_TTFTRow] | None] = await prisma_client.db.query_raw( + db_response: Final[Sequence[_TTFTRow] | None] = await prisma_client.replica_db.query_raw( sql_query, _selected_model_group, startTime, endTime ) _daily_entries: dict = {} # {"Jun 23": {"model1": 0.002, "model2": 0.003}} @@ -15267,7 +15267,7 @@ async def model_metrics( avg_latency_per_token DESC; """ _all_api_bases: Final = set() - db_response: Final[Sequence[_LatencyRow] | None] = await prisma_client.db.query_raw( + db_response: Final[Sequence[_LatencyRow] | None] = await prisma_client.replica_db.query_raw( sql_query, _selected_model_group, startTime, endTime, api_key, customer ) _daily_entries: dict = {} # {"Jun 23": {"model1": 0.002, "model2": 0.003}} @@ -15385,7 +15385,7 @@ ORDER BY slow_count DESC; """ - db_response: Final = await prisma_client.db.query_raw( + db_response: Final = await prisma_client.replica_db.query_raw( sql_query, alerting_threshold, _selected_model_group, @@ -15458,7 +15458,7 @@ async def model_metrics_exceptions( ORDER BY total_exceptions DESC LIMIT 200; """ - db_response: Final[Sequence[_ExceptionRow] | None] = await prisma_client.db.query_raw( + db_response: Final[Sequence[_ExceptionRow] | None] = await prisma_client.replica_db.query_raw( sql_query, startTime, endTime, _selected_model_group, api_key ) response: Final[list[dict]] = [] diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py index 376b113ed02..9515fecb586 100644 --- a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -177,7 +177,7 @@ async def _advance_marker(prisma_client: "PrismaClient", days: tuple[str, ...], async def _db_now(prisma_client: "PrismaClient") -> _NowRow: - rows: Final = await prisma_client.db.query_raw(_DB_NOW_SQL) + rows: Final = await prisma_client.replica_db.query_raw(_DB_NOW_SQL) return _NowRow.model_validate(rows[0]) @@ -189,9 +189,9 @@ async def _scan_pending(prisma_client: "PrismaClient") -> _PendingScan: db_now: Final = await _db_now(prisma_client) last_closed_day: Final = (date.fromisoformat(db_now.today) - timedelta(days=1)).isoformat() rows: Final = ( - await prisma_client.db.query_raw(_ALL_CLOSED_DAYS_SQL, last_closed_day) + await prisma_client.replica_db.query_raw(_ALL_CLOSED_DAYS_SQL, last_closed_day) if marker is None or marker.scanned_at is None - else await prisma_client.db.query_raw( + else await prisma_client.replica_db.query_raw( _PENDING_DAYS_SQL, last_closed_day, marker.reconciled_through, marker.scanned_at ) ) diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index ee2e1cfeaf7..a88071d6128 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -134,7 +134,7 @@ async def _reverse_hash_key_metadata( warning: str, ) -> Mapping[str, KeyMetadataDict]: rows: Final = await _db_or_empty( - lambda: prisma_client.db.query_raw(sql, sorted(wanted)), + lambda: prisma_client.replica_db.query_raw(sql, sorted(wanted)), warning, len(wanted), ) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 822b827f985..d5c612a98f3 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -233,7 +233,7 @@ class _SpendDailySummaryRow(TypedDict): async def _query_raw(prisma_client: PrismaClient, sql_query: str, *args: object) -> Sequence[_RowT]: """Run a raw read query and return its rows as the row type the caller declares.""" - return await prisma_client.db.query_raw(sql_query, *args) + return await prisma_client.replica_db.query_raw(sql_query, *args) async def _query_raw_or_none(prisma_client: PrismaClient, sql_query: str, *args: object) -> Sequence[_RowT] | None: @@ -2922,7 +2922,7 @@ async def ui_view_spend_logs( ) sql_params.extend([page_size, skip]) - data: Final = await prisma_client.db.query_raw(sql_query, *sql_params) + data: Final = await prisma_client.replica_db.query_raw(sql_query, *sql_params) if request_id is not None and not is_v2 and not is_admin_view: await _assert_user_owns_fetched_spend_rows( diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 8f85ecdd480..bf904d58b9b 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -899,7 +899,7 @@ async def _query_raw_rows( sql_query: str, *args: object, ) -> Sequence[Mapping[str, object]] | None: - return await prisma_client.db.query_raw(sql_query, *args) + return await prisma_client.replica_db.query_raw(sql_query, *args) async def get_spend_by_team( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index bc64293c9b3..6a20dc3ffdd 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4511,6 +4511,11 @@ class PrismaClient: return self.db.read_target return self.db + @property + def replica_db(self) -> "PrismaWrapper | RoutingPrismaWrapper": + """Explicit opt-in to read-replica routing; reads reached through it may go to the reader.""" + return self.db + def tx(self, *, timeout: timedelta = _PRISMA_DEFAULT_TX_TIMEOUT) -> "TransactionManager": """Open an interactive transaction on the writer. @@ -4597,7 +4602,7 @@ class PrismaClient: required_view: Final = "LiteLLM_VerificationTokenView" expected_views_str: Final = ", ".join(f"'{view}'" for view in expected_views) pg_schema: Final = os.getenv("DATABASE_SCHEMA", "public") - ret: Final[Sequence[_ViewCountRow]] = await self.db.query_raw(f""" + ret: Final[Sequence[_ViewCountRow]] = await self.replica_db.query_raw(f""" WITH existing_views AS ( SELECT viewname FROM pg_views @@ -4636,9 +4641,9 @@ class PrismaClient: """, ) else: - should_create_views: Final = await should_create_missing_views(db=self.db) + should_create_views: Final = await should_create_missing_views(db=self.replica_db) if should_create_views: - await create_missing_views(db=self.db) + await create_missing_views(db=self.replica_db) else: # don't block execution if these views are missing # Convert lists to sets for efficient difference calculation @@ -4692,7 +4697,7 @@ class PrismaClient: ) return await config_table.find_first(where={key: value}) elif table_name == "spend": - return await self.db.l.find_first(where={key: value}) + return await self.replica_db.l.find_first(where={key: value}) return None try: @@ -4764,7 +4769,7 @@ class PrismaClient: """ stale_read_engine: Final = _StaleReadEngine.observe(self.read_db) try: - return await self.db.query_first(sql_query, *args) + return await self.replica_db.query_first(sql_query, *args) except Exception as e: if "cached plan must not change result type" not in str(e): raise @@ -4779,7 +4784,7 @@ class PrismaClient: force_recreate=True, stale_read_engine=stale_read_engine, ) - return await self.db.query_first(sql_query, *args) + return await self.replica_db.query_first(sql_query, *args) @backoff.on_exception( backoff.expo, @@ -4967,7 +4972,7 @@ class PrismaClient: LIMIT $1 OFFSET $2 """ - response = await self.db.query_raw(sql_query, limit, offset) + response = await self.replica_db.query_raw(sql_query, limit, offset) return response elif table_name == "spend": verbose_proxy_logger.debug("PrismaClient: get_data: table_name == 'spend'") @@ -5113,7 +5118,9 @@ class PrismaClient: # If not found in main table, check deprecated keys (grace period) # check_deprecated=False on the recursive call prevents unbounded chaining if response is None and hashed_token is not None and check_deprecated: - active_token_id: Final = await _lookup_deprecated_key(db=self.db, hashed_token=hashed_token) + active_token_id: Final = await _lookup_deprecated_key( + db=self.replica_db, hashed_token=hashed_token + ) if active_token_id: # The recursive call returns a finished # LiteLLM_VerificationTokenView; the dict @@ -6557,7 +6564,7 @@ class PrismaClient: async def _view_setup_gate_table_present(self) -> bool: rows: Final = _VIEW_SETUP_GATE_PROBE_ROWS.validate_python( - await self.db.query_raw("SELECT to_regclass($1) IS NOT NULL AS present", _VIEW_SETUP_GATE_TABLE) + await self.replica_db.query_raw("SELECT to_regclass($1) IS NOT NULL AS present", _VIEW_SETUP_GATE_TABLE) ) return rows[0]["present"] @@ -6566,7 +6573,7 @@ class PrismaClient: try: await asyncio.sleep(self._db_health_watchdog_interval_seconds) await asyncio.wait_for( - self.db.query_raw("SELECT 1"), + self.replica_db.query_raw("SELECT 1"), timeout=self._db_health_watchdog_probe_timeout_seconds, ) if isinstance(self.db, RoutingPrismaWrapper) and self.db.writer_unavailable: @@ -6778,7 +6785,7 @@ class PrismaClient: FROM pg_class WHERE oid = '"LiteLLM_SpendLogs"'::regclass; """ - result: Final[Sequence[_RelTuplesRow]] = await self.db.query_raw(query=sql_query) + result: Final[Sequence[_RelTuplesRow]] = await self.replica_db.query_raw(query=sql_query) return result[0]["reltuples"] try: diff --git a/litellm/repositories/autorouter_session_repository.py b/litellm/repositories/autorouter_session_repository.py index d05ef9421ca..63479d6ca51 100644 --- a/litellm/repositories/autorouter_session_repository.py +++ b/litellm/repositories/autorouter_session_repository.py @@ -15,7 +15,7 @@ if TYPE_CHECKING: class AutoRouterSessionRepository(BaseRepository[LiteLLM_AutoRouterSession]): @property def table(self) -> TableActions["prisma_models.LiteLLM_AutoRouterSession"]: - return self.prisma_client.db.litellm_autoroutersession + return self.prisma_client.replica_db.litellm_autoroutersession @property def model_class(self) -> type[LiteLLM_AutoRouterSession]: diff --git a/litellm/repositories/budget_repository.py b/litellm/repositories/budget_repository.py index 205646c8393..89fb64fd52f 100644 --- a/litellm/repositories/budget_repository.py +++ b/litellm/repositories/budget_repository.py @@ -14,7 +14,7 @@ if TYPE_CHECKING: class _BudgetDb(Protocol): - """The single Prisma table this repository reaches for on ``prisma_client.db``.""" + """The single Prisma table this repository reaches for on ``prisma_client.replica_db``.""" @property def litellm_budgettable(self) -> TableActions["prisma_models.LiteLLM_BudgetTable"]: ... @@ -24,7 +24,7 @@ class _PrismaClientView(Protocol): """The one attribute this repository reads off the untyped Prisma client wrapper.""" @property - def db(self) -> _BudgetDb: ... + def replica_db(self) -> _BudgetDb: ... class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): @@ -33,7 +33,7 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_BudgetTable"]: client: Final[_PrismaClientView] = self.prisma_client - return client.db.litellm_budgettable + return client.replica_db.litellm_budgettable @property def model_class(self) -> type[LiteLLM_BudgetTable]: diff --git a/litellm/repositories/config_repository.py b/litellm/repositories/config_repository.py index 8b8280622fd..ff1dae91b69 100644 --- a/litellm/repositories/config_repository.py +++ b/litellm/repositories/config_repository.py @@ -55,7 +55,7 @@ class ConfigRepository: @property def _config_table(self) -> _ConfigTable: - return cast(_ConfigTable, self.prisma_client.db.litellm_config) + return cast(_ConfigTable, self.prisma_client.replica_db.litellm_config) @property def table(self) -> _ConfigTable: diff --git a/litellm/repositories/credentials_repository.py b/litellm/repositories/credentials_repository.py index ddb9767b2b9..41eec940196 100644 --- a/litellm/repositories/credentials_repository.py +++ b/litellm/repositories/credentials_repository.py @@ -27,7 +27,7 @@ class _PrismaCredentialsDb(Protocol): class _PrismaClientView(Protocol): @property - def db(self) -> _PrismaCredentialsDb: ... + def replica_db(self) -> _PrismaCredentialsDb: ... class CredentialsRepository: @@ -46,7 +46,7 @@ class CredentialsRepository: @property def table(self) -> "_CredentialsTable": return wrap_table_actions_for_config_sync( - actions=self.prisma_client.db.litellm_credentialstable, + actions=self.prisma_client.replica_db.litellm_credentialstable, table_name="litellm_credentialstable", ) diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 7736939c696..d5d5a0044f2 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -17,7 +17,7 @@ class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]: - return self.prisma_client.db.litellm_objectpermissiontable + return self.prisma_client.replica_db.litellm_objectpermissiontable @property def model_class(self) -> type[LiteLLM_ObjectPermissionTable]: diff --git a/litellm/repositories/organization_repository.py b/litellm/repositories/organization_repository.py index 47eb8f4a609..596a682a97e 100644 --- a/litellm/repositories/organization_repository.py +++ b/litellm/repositories/organization_repository.py @@ -14,7 +14,7 @@ if TYPE_CHECKING: class _OrganizationDb(Protocol): - """The single Prisma table this repository reaches for on ``prisma_client.db``.""" + """The single Prisma table this repository reaches for on ``prisma_client.replica_db``.""" @property def litellm_organizationtable(self) -> TableActions["prisma_models.LiteLLM_OrganizationTable"]: ... @@ -24,7 +24,7 @@ class _PrismaClientView(Protocol): """The one attribute this repository reads off the untyped Prisma client wrapper.""" @property - def db(self) -> _OrganizationDb: ... + def replica_db(self) -> _OrganizationDb: ... class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): @@ -33,7 +33,7 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_OrganizationTable"]: client: Final[_PrismaClientView] = self.prisma_client - return client.db.litellm_organizationtable + return client.replica_db.litellm_organizationtable @property def model_class(self) -> type[LiteLLM_OrganizationTable]: diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 60c16fbd746..55736300b9a 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -1,7 +1,7 @@ """ Typed Protocol seams over prisma-client-py surfaces. -Modules that reach Prisma through an untyped handle (``prisma_client.db`` or a +Modules that reach Prisma through an untyped handle (``prisma_client.replica_db`` or a repository ``.table``) annotate against these Protocols instead of hand-rolling private ones per file. """ @@ -14,7 +14,7 @@ RowT_co = TypeVar("RowT_co", covariant=True) class DatabaseClient(Protocol): @property - def db(self) -> object: ... + def replica_db(self) -> object: ... class TableActions(Protocol[RowT_co]): diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 905e813f35e..639f520649b 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -18,7 +18,7 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: - return self.prisma_client.db.litellm_projecttable + return self.prisma_client.replica_db.litellm_projecttable @property def model_class(self) -> type[LiteLLM_ProjectTable]: diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 1ad7a735d96..ef95293682e 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -32,7 +32,7 @@ class PrismaTableRepository(Generic[RowT_co]): @property def table(self) -> TableActions[RowT_co]: - actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.db, self.table_name) + actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.replica_db, self.table_name) return wrap_table_actions_for_config_sync(actions=actions, table_name=self.table_name) diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index cbe263699c9..9a6d0b2203a 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -68,7 +68,7 @@ class _PrismaTeamDb(_TeamTables, Protocol): class _PrismaClientView(Protocol): @property - def db(self) -> _PrismaTeamDb: ... + def replica_db(self) -> _PrismaTeamDb: ... _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) @@ -88,7 +88,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): @property def _db(self) -> _PrismaTeamDb: client: Final[_PrismaClientView] = self.prisma_client - return client.db + return client.replica_db @property def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]: diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index 87eb45f262d..66ed364039e 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -40,7 +40,7 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]: - return self.prisma_client.db.litellm_usertable + return self.prisma_client.replica_db.litellm_usertable @property def model_class(self) -> type[LiteLLM_UserTable]: diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index d02c2114136..87f86fa81d4 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -54,7 +54,7 @@ class _PrismaVerificationTokenDb(_VerificationTokenTables, Protocol): class _PrismaClientView(Protocol): @property - def db(self) -> _PrismaVerificationTokenDb: ... + def replica_db(self) -> _PrismaVerificationTokenDb: ... _JSON_ENCODED_TOKEN_FIELDS: Final = ( @@ -76,7 +76,7 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): @property def _db(self) -> _PrismaVerificationTokenDb: client: Final[_PrismaClientView] = self.prisma_client - return client.db + return client.replica_db @property def table(self) -> TableActions["PrismaVerificationToken"]: diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index f749977eb82..db7a4f88376 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -308,7 +308,7 @@ class ResponsesSessionHandler: for attempt in range(max_attempts): if attempt: await asyncio.sleep(RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL) - if spend_logs := await prisma_client.db.query_raw(query, response_id): + if spend_logs := await prisma_client.replica_db.query_raw(query, response_id): verbose_proxy_logger.debug( "Found the following spend logs for previous response id %s: %s", response_id, diff --git a/tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py b/tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py index 45e1acecec8..18e00336130 100644 --- a/tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py +++ b/tests/test_litellm/integrations/SlackAlerting/test_user_spend_alerts.py @@ -152,6 +152,7 @@ async def test_send_user_spend_alerts_sends_and_dedupes(): alerting_args={"daily_spend_per_user_threshold": 50.0, "spend_anomaly_min_spend": 1000.0}, ) mock_prisma: Final = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( return_value=[ { @@ -189,5 +190,6 @@ async def test_send_user_spend_alerts_noop_when_alert_types_disabled(): alerting_args={"daily_spend_per_user_threshold": 50.0}, ) mock_prisma: Final = AsyncMock() + mock_prisma.replica_db = mock_prisma.db await slack_alerting.send_user_spend_alerts(prisma_client=mock_prisma) mock_prisma.db.query_raw.assert_not_called() diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/test_litellm/integrations/test_shadow_eval_logger.py index e3f059a7941..776936f4975 100644 --- a/tests/test_litellm/integrations/test_shadow_eval_logger.py +++ b/tests/test_litellm/integrations/test_shadow_eval_logger.py @@ -62,6 +62,7 @@ def _job(**overrides) -> ActiveShadowEvalJob: def _prisma(jobs=(), attempt_counts=(), attempt_costs=()) -> MagicMock: costs = {job_id: {"judge_cost": judge, "shadow_cost": shadow} for job_id, judge, shadow in attempt_costs} prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=list(jobs)) prisma.db.litellm_shadowevalattempt.group_by = AsyncMock( return_value=[ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py index 5d6d47b8c38..68b17959264 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py @@ -60,6 +60,7 @@ def _make_prisma(stored: dict, db_has_id_jag_server: bool = False): ``db_has_id_jag_server`` drives the retention gate's authoritative DB fallback; it is wired explicitly so the gate never reads a truthy bare MagicMock.""" prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_mcpservertable.find_first = AsyncMock(return_value=MagicMock() if db_has_id_jag_server else None) async def _upsert(where, data): @@ -417,6 +418,7 @@ async def test_retain_none_assertion_never_consults_gate_or_store(): @pytest.mark.asyncio async def test_retain_swallows_store_failure(): prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=RuntimeError("db down")) with ( patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, @@ -465,6 +467,7 @@ async def test_db_store_converts_a_driver_failure_into_assertion_store_unavailab """The live store must not let a raw driver error escape: the resolver distinguishes an outage from an absent assertion, and only a typed failure lets it do that.""" prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_ssoidentityassertion.find_unique = AsyncMock(side_effect=RuntimeError("connection refused")) with patch("litellm.proxy.proxy_server.prisma_client", prisma): with pytest.raises(AssertionStoreUnavailable): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py index c1239c228aa..6b6eb2c010c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py @@ -133,6 +133,7 @@ def _byok_key_row(server_id): def _mock_prisma(null_rows, token_rows): mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=null_rows) mock_prisma.db.litellm_mcpservertable.update_many = AsyncMock(return_value=MagicMock()) mock_prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=token_rows) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py index b6c946b95fa..8cbc2542d7c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py @@ -28,6 +28,7 @@ def _row(**overrides): def _prisma(rows): prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows) prisma_client.db.litellm_mcpservertable.update = AsyncMock() return prisma_client diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index 482294e7b92..21d522495cf 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -76,6 +76,7 @@ client = TestClient(app) @pytest.fixture def mock_prisma_client(): with patch("litellm.proxy.proxy_server.prisma_client") as mock: + mock.replica_db = mock.db yield mock @@ -208,6 +209,7 @@ def test_agent_error_schema_consistency( @pytest.mark.asyncio async def test_get_agent_daily_activity_admin_param_passing(monkeypatch): mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) @@ -246,6 +248,7 @@ async def test_get_agent_daily_activity_admin_param_passing(monkeypatch): @pytest.mark.asyncio async def test_get_agent_daily_activity_with_agent_names(monkeypatch): mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_agent1 = MagicMock() mock_agent1.agent_id = "agent-1" mock_agent1.agent_name = "First Agent" @@ -303,6 +306,7 @@ async def test_attach_keys_to_agents_groups_by_agent_and_omits_secret(): agent_without_keys = _sample_agent_response(agent_id="agent-2") mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[ _Row("hash-aaa", "agent-1", "primary", "sk-...aaa"), @@ -348,6 +352,7 @@ class TestAgentByIdKeyRedaction: test_client = _make_app_with_role(role) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value=None ) @@ -410,6 +415,7 @@ class TestAgentRBACInternalUser: return_value=_sample_agent_response() ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value=None ) @@ -532,6 +538,7 @@ class TestAgentRBACProxyAdminViewOnly: key_row.key_name = "sk-...aaa" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[key_row] @@ -578,6 +585,7 @@ class TestAgentRBACProxyAdmin: def test_should_allow_admin_to_create_agent(self, monkeypatch): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db self.mock_registry.get_agent_by_name = MagicMock(return_value=None) self.mock_registry.add_agent_to_db = AsyncMock( return_value=_sample_agent_response() @@ -660,6 +668,7 @@ class TestAgentRBACProxyAdmin: def test_update_agent_response_never_echoes_secret(self): """LIT-6736: PUT /v1/agents/{id} must not echo the stored secret back.""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value={ "agent_id": "agent-123", @@ -695,6 +704,7 @@ class TestAgentRBACProxyAdmin: def test_patch_agent_response_never_echoes_secret(self): """LIT-6736: PATCH /v1/agents/{id} must not echo the stored secret back.""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: # test-quality-ok: proxy_server module global is the endpoint's only injection point + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value={ "agent_id": "agent-123", @@ -730,6 +740,7 @@ class TestAgentRBACProxyAdmin: "agent_card_params": _sample_agent_card_params(), } with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( return_value=existing ) diff --git a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py index 0fd0dda3017..d5c50f6b03f 100644 --- a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py +++ b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py @@ -117,6 +117,7 @@ class CountingRedis(RedisCache): def _prisma(rows: dict[str, object] | None = ALL_ROWS) -> MagicMock: prisma = MagicMock(name="prisma_client") + prisma.replica_db = prisma.db prisma.db.query_first = AsyncMock(return_value=rows) return prisma diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index f03abe8f124..46175e73b2c 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -2049,6 +2049,7 @@ async def test_auto_register_binds_api_key_to_token_hash(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() user_api_key_cache = MagicMock() @@ -2105,6 +2106,7 @@ async def test_auto_register_first_request_propagates_user_email(active: bool) - general_settings = {"enable_jwt_auth": True} user_api_key_cache = DualCache() prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db jwt_handler = MagicMock() jwt_handler.is_jwt.return_value = True jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"}) @@ -2223,6 +2225,7 @@ async def test_auto_register_stamps_new_key_with_jwt_agent_id(): credential_ref=CredentialRef(token_id=token_hash), ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.create = AsyncMock() user_api_key_cache = MagicMock() user_api_key_cache.async_set_cache = AsyncMock() @@ -2280,6 +2283,7 @@ async def test_auto_register_race_loser_keeps_winners_agent_id(losing_agent_id: credential_ref=CredentialRef(token_id=winner_hash), ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_jwtkeymapping.create = AsyncMock( side_effect=Exception("Unique constraint failed on the fields: (`jwt_claim_name`,`jwt_claim_value`)") ) @@ -2661,6 +2665,7 @@ class TestJWTOAuth2Coexistence: general_settings = {"enable_jwt_auth": True} user_api_key_cache = DualCache() prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db jwt_handler = MagicMock() jwt_handler.is_jwt.return_value = True jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "user1"}) @@ -2759,6 +2764,7 @@ class TestJWTOAuth2Coexistence: general_settings = {"enable_jwt_auth": True} user_api_key_cache = DualCache() prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db jwt_handler = MagicMock() jwt_handler.is_jwt.return_value = True jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "mapped-user"}) @@ -2837,6 +2843,7 @@ class TestJWTOAuth2Coexistence: general_settings = {"enable_jwt_auth": True} user_api_key_cache = DualCache() prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db jwt_handler = MagicMock() jwt_handler.is_jwt.return_value = True jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "jwt-principal"}) @@ -2914,6 +2921,7 @@ class TestJWTOAuth2Coexistence: general_settings = {"enable_jwt_auth": True} user_api_key_cache = DualCache() prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db jwt_handler = MagicMock() jwt_handler.is_jwt.return_value = True jwt_handler.auth_jwt = AsyncMock(return_value={"sub": "mapped-user"}) @@ -4421,6 +4429,7 @@ async def _run_centralized_checks_with_key_end_user_budget( return _end_user_budget_row(budget_id, budgets[budget_id]) if budget_id in budgets else None prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[]) prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row) prisma_client.db.litellm_budgettable.find_unique = AsyncMock(side_effect=_find_budget) @@ -4687,6 +4696,7 @@ def _unrestricted_end_user_prisma(spend: float): end_user_row.dict = lambda: {"user_id": "customer-1", "blocked": False, "spend": spend} mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row) mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) @@ -6661,8 +6671,10 @@ def _proxy_attrs_for_db_lookup(): ``_user_api_key_auth_builder`` down to the DB key lookup.""" proxy_logging_obj = MagicMock() proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db return { - "prisma_client": MagicMock(), + "prisma_client": prisma_client, "user_api_key_cache": DualCache(), "proxy_logging_obj": proxy_logging_obj, "master_key": "sk-test-master", @@ -7735,6 +7747,7 @@ async def test_global_proxy_spend_reads_resettable_proxy_budget_row(): proxy_budget_row = MagicMock() proxy_budget_row.spend = 42.5 prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=proxy_budget_row) prisma_client.db.query_raw = AsyncMock( side_effect=AssertionError("global spend must not be loaded from the fixed-30d MonthlyGlobalSpend view") @@ -7762,6 +7775,7 @@ async def test_global_proxy_spend_none_when_proxy_budget_row_missing(): from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) result = await _fetch_global_spend_with_event_coordination( @@ -8516,7 +8530,7 @@ def _per_issuer_virtual_key_jwt_handler( def _fake_prisma_with_jwt_key_mapping(hashed_token: str | None) -> tuple[SimpleNamespace, AsyncMock]: """Every ``find_first`` call (issuer-scoped or global fallback) resolves the same way.""" find_first = AsyncMock(return_value=None if hashed_token is None else SimpleNamespace(token=hashed_token)) - prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first))) + prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)), replica_db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first))) return prisma_client, find_first @@ -8535,7 +8549,7 @@ def _fake_prisma_jwt_key_mapping_table(rows: list[dict[str, object]]) -> tuple[S return None find_first = AsyncMock(side_effect=_find_first) - prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first))) + prisma_client = SimpleNamespace(db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first)), replica_db=SimpleNamespace(litellm_jwtkeymapping=SimpleNamespace(find_first=find_first))) return prisma_client, find_first diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 131db55ee01..58bad1e050f 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -156,6 +156,7 @@ class MockPrismaClient: } self.get_data_calls: List[Dict[str, Any]] = [] self.db = MockDB() + self.replica_db = self.db async def get_data(self, table_name, query_type, **kwargs): self.get_data_calls.append({"table_name": table_name, "query_type": query_type, **kwargs}) @@ -933,6 +934,7 @@ def _make_reset_budget_windows_job( Returns (job, prisma_client_mock, spend_counter_cache_mock). """ prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db async def fake_query_raw(query: str, *args, **kwargs): # Dispatch by table name in the SQL so a single stub covers both calls. @@ -1239,6 +1241,7 @@ def test_reset_budget_windows_query_error_does_not_break_team_path(monkeypatch): expired = (now - timedelta(minutes=1)).isoformat() + "Z" prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db async def fake_query_raw(query: str, *args, **kwargs): if '"LiteLLM_VerificationToken"' in query: @@ -1468,6 +1471,7 @@ def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch): now = datetime.now(timezone.utc) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db matching_key = type( "Key", @@ -1940,6 +1944,7 @@ def _job_with_expired_budget(db, proxy_logging=None): something to invalidate and its absence is a real signal.""" prisma_client = MockPrismaClient() prisma_client.db = db + prisma_client.replica_db = prisma_client.db prisma_client.data["budget"] = [_budget_row(budget_id="budget-1", budget_duration="7d")] db.litellm_tagtable.set_find_many_results([type("Tag", (), {"tag_name": "tenant-42"})]) job = ResetBudgetJob( @@ -2344,6 +2349,7 @@ def test_budget_table_reset_stops_when_the_cascade_fails(monkeypatch): monkeypatch.setattr(reset_budget_job_module, "RESET_BUDGET_JOB_BATCH_SIZE", 2) client, job = _chunked_job({"budget": [[_budget_row("b1"), _budget_row("b2")]]}) client.db = FailingCommitDB() + client.replica_db = client.db asyncio.run(job.reset_budget_for_litellm_budget_table()) @@ -2620,6 +2626,7 @@ def _paginating_window_job(monkeypatch, pages_by_table: Dict[str, List[List[Dict Returns (job, calls) where calls is a list of (sql, cursor, limit). """ prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db remaining = {table: list(pages) for table, pages in pages_by_table.items()} calls: List[Dict[str, Any]] = [] @@ -2696,6 +2703,7 @@ def test_reset_budget_windows_pages_to_the_end_of_a_large_table(monkeypatch): def test_reset_budget_windows_survives_one_table_failing(monkeypatch): """A broken key scan must not cost the team scan its sweep.""" prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db async def fake_query_raw(query: str, *args, **kwargs): if '"LiteLLM_VerificationToken"' in query: @@ -2759,6 +2767,7 @@ def _cursor_paginating_window_job(monkeypatch, key_rows: List[Dict[str, Any]]): genuinely re-reads the same prefix. """ prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db ordered = sorted(key_rows, key=lambda r: r["token"]) visited: List[str] = [] diff --git a/tests/test_litellm/proxy/db/test_budget_window_spend_writer.py b/tests/test_litellm/proxy/db/test_budget_window_spend_writer.py index 130f0c56ccf..2480498364b 100644 --- a/tests/test_litellm/proxy/db/test_budget_window_spend_writer.py +++ b/tests/test_litellm/proxy/db/test_budget_window_spend_writer.py @@ -68,6 +68,7 @@ class _FakeDB: class _FakePrismaClient: def __init__(self, db: _FakeDB) -> None: self.db = db + self.replica_db = self.db class _RecordingAggregate: diff --git a/tests/test_litellm/proxy/db/test_health_check_latest.py b/tests/test_litellm/proxy/db/test_health_check_latest.py index 6322891ae9e..f4bffc3a9ac 100644 --- a/tests/test_litellm/proxy/db/test_health_check_latest.py +++ b/tests/test_litellm/proxy/db/test_health_check_latest.py @@ -14,6 +14,7 @@ from litellm.proxy.db.health_check_latest import ( def _prisma(rows): prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.query_raw = AsyncMock(return_value=rows) return prisma diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py index 2b438a9d370..8e3531b6a8a 100644 --- a/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_budgets.py @@ -76,6 +76,7 @@ def query_raw(monkeypatch): """Mocks the one call the executor makes. `count` reads the first result, `find_many` the second.""" mock = AsyncMock(side_effect=[[{"count": 0}], []]) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.query_raw = mock monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) return mock diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_spend_logs.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_spend_logs.py index b6867d338c5..2f1268e0162 100644 --- a/tests/test_litellm/proxy/management_endpoints/management_v1/test_spend_logs.py +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_spend_logs.py @@ -51,6 +51,7 @@ WINDOW = "filter[startTime][gte]=2026-07-23T00:00:00Z&filter[startTime][lte]=202 @pytest.fixture def mock_prisma_client(monkeypatch): prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.query_raw = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) return prisma_client diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index ff3d19e8637..b057fe386fe 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -643,7 +643,8 @@ class TestAutoRouterBenchmarks: async def query_raw(self, sql: str, *params: object): return rows - monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})()) + db = _DB() + monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": db, "replica_db": db})()) monkeypatch.setattr(proxy_server, "llm_router", type("R", (), {"model_list": model_list})()) return await get_auto_router_benchmarks( user_api_key_dict=ADMIN, @@ -817,7 +818,7 @@ class TestAutoRouterBenchmarks: from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_benchmarks query: Final = AsyncMock(return_value=[]) - monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query))) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(query_raw=query), replica_db=SimpleNamespace(query_raw=query))) app: Final = FastAPI() app.get("/auto_router/benchmarks")(get_auto_router_benchmarks) app.dependency_overrides[user_api_key_auth] = lambda: ADMIN @@ -858,7 +859,8 @@ class TestAutoRouterBenchmarks: captured["params"] = params return [TestAutoRouterBenchmarks.ROW.model_dump()] - monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})()) + db = _DB() + monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": db, "replica_db": db})()) response = await get_auto_router_benchmarks( user_api_key_dict=UserAPIKeyAuth(user_role=role, api_key="sk-admin", user_id="viewer"), @@ -921,7 +923,8 @@ class TestAutoRouterBenchmarks: async def query_raw(self, sql: str, *params: object): return [{**TestAutoRouterBenchmarks.ROW.model_dump(), "tier_turns": wire_value}] - monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": _DB()})()) + db = _DB() + monkeypatch.setattr(proxy_server, "prisma_client", type("P", (), {"db": db, "replica_db": db})()) response = await get_auto_router_benchmarks( user_api_key_dict=ADMIN, @@ -1091,10 +1094,11 @@ class TestAutoRouterSession: ] return max(matching, key=lambda r: r["last_turn_at"], default=None) + db = type("D", (), {"litellm_autoroutersession": _Table()})() monkeypatch.setattr( proxy_server, "prisma_client", - type("P", (), {"db": type("D", (), {"litellm_autoroutersession": _Table()})()})(), + type("P", (), {"db": db, "replica_db": db})(), ) return lookups @@ -1354,6 +1358,7 @@ def _shadow_prisma( direction sees the opposite-direction legs a key may hold at the same time, and a group read that matched on a leg id would come back empty.""" prisma = MagicMock() + prisma.replica_db = prisma.db teams: Final = key_teams or {} team_aliases: Final = known_teams or {} user_emails: Final = known_users or {} @@ -3085,6 +3090,7 @@ async def test_routing_test_never_confirms_models_the_caller_cannot_use(monkeypa team_row.model_dump.return_value = row_data team_row.dict.return_value = row_data prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) return prisma @@ -3146,6 +3152,7 @@ async def test_validate_config_gates_like_the_write_it_rehearses(monkeypatch: py "members_with_roles": [{"role": "admin", "user_id": "team-admin"}], } prisma: Final = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) monkeypatch.setattr(proxy_server, "prisma_client", prisma) monkeypatch.setattr(proxy_server, "premium_user", True) @@ -3178,6 +3185,7 @@ def _configure_member_preview(monkeypatch: pytest.MonkeyPatch, *, allowed: bool team_member_permissions=["/auto_router/manage"] if allowed else [], ) prisma: Final = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) monkeypatch.setattr(proxy_server, "prisma_client", prisma) @@ -3615,7 +3623,7 @@ async def test_availability_counts_db_and_yaml_without_disclosing_router_names(m monkeypatch.setattr( proxy_server, "prisma_client", - SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))), + SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many)), replica_db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))), ) monkeypatch.setattr(proxy_server.proxy_config, "auto_router_db_catalog", build_auto_router_catalog((row,))) monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(config_deployments=lambda: (yaml_row,))) @@ -3659,7 +3667,7 @@ async def test_availability_denies_another_teams_edit_exemption(monkeypatch): monkeypatch.setattr( proxy_server, "prisma_client", - SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))), + SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many)), replica_db=SimpleNamespace(litellm_proxymodeltable=SimpleNamespace(find_many=find_many))), ) monkeypatch.setattr(proxy_server.proxy_config, "auto_router_db_catalog", build_auto_router_catalog((row,))) monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(config_deployments=lambda: ())) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index baaf3f4ba2f..bd98da99996 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -42,6 +42,7 @@ async def test_get_daily_activity_empty_entity_id_list(): # Mock PrismaClient mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db # Mock the table methods mock_table = MagicMock() @@ -94,6 +95,7 @@ async def test_get_daily_activity_order_has_id_tiebreaker(): """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_table = MagicMock() mock_table.count = AsyncMock(return_value=0) mock_table.find_many = AsyncMock(return_value=[]) @@ -147,6 +149,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): # Mock PrismaClient mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db # query_raw now returns rollup rows produced by GROUPING SETS, each # tagged with its grouping level via GROUPING_ID(). The dispatcher @@ -312,6 +315,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): async def test_get_api_key_metadata_returns_active_key_metadata(): """Test that get_api_key_metadata should return metadata for active keys.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # Mock active key record mock_active_key = MagicMock() @@ -335,6 +339,7 @@ async def test_get_api_key_metadata_returns_active_key_metadata(): async def test_get_api_key_metadata_falls_back_to_deleted_keys(): """Test that get_api_key_metadata should fall back to deleted keys table for missing keys.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # No active keys found mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -367,6 +372,7 @@ async def test_get_api_key_metadata_falls_back_to_deleted_keys(): async def test_get_api_key_metadata_mixed_active_and_deleted_keys(): """Test that get_api_key_metadata should return metadata for both active and deleted keys.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # One active key found mock_active_key = MagicMock() @@ -401,6 +407,7 @@ async def test_get_api_key_metadata_mixed_active_and_deleted_keys(): async def test_get_api_key_metadata_deleted_table_not_queried_when_all_keys_found(): """Test that get_api_key_metadata should not query deleted table when all keys are active.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_active_key = MagicMock() mock_active_key.token = "key-hash-1" @@ -426,6 +433,7 @@ async def test_get_api_key_metadata_deleted_table_not_queried_when_all_keys_foun async def test_get_api_key_metadata_deleted_table_error_handled_gracefully(): """Test that get_api_key_metadata should handle errors from deleted table gracefully.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # No active keys found mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -446,6 +454,7 @@ async def test_get_api_key_metadata_deleted_table_error_handled_gracefully(): async def test_get_api_key_metadata_regenerated_key_uses_most_recent_deleted_record(): """Test that get_api_key_metadata should use the most recent deleted record for regenerated keys.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # No active keys found (old hash no longer in active table after regeneration) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -486,6 +495,7 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash( double_hashed = hash_token("a" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock( @@ -515,6 +525,7 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s double_hashed = hash_token("b" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) @@ -563,6 +574,7 @@ async def test_get_api_key_metadata_permanent_miss_with_a_window_reads_spend_log double_hashed = hash_token("permanent-miss-with-window-6852") window = (datetime(2024, 1, 1), datetime(2024, 1, 4)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) @@ -586,6 +598,7 @@ async def test_get_daily_activity_recovers_a_session_key_alias_from_spend_logs_a records = [_daily_user_spend_record(user_id="session-user", api_key=session_digest, spend=1.5)] mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_table = MagicMock() mock_table.count = AsyncMock(return_value=len(records)) mock_table.find_many = AsyncMock(return_value=records) @@ -738,6 +751,7 @@ async def test_tag_daily_activity_metadata_totals_not_zero(): """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db # Create mock tag spend records (request_id is NULL for aggregated rows) mock_record_1 = MagicMock() @@ -841,6 +855,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): """Test that the full aggregation pipeline should preserve metadata for deleted keys.""" mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db # GROUPING SETS rollup rows. The api_key metadata lookup is driven # by any non-NULL api_key in the result set, so the (date, endpoint, @@ -935,6 +950,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): async def test_aggregated_activity_flags_only_keys_that_key_info_can_still_resolve(): """/key/info reads the active key table only, so deleted and never-stored (session) keys must not claim to exist.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db base = { "date": "2024-01-01", "endpoint": "/v1/chat/completions", @@ -1033,6 +1049,7 @@ async def test_get_daily_activity_applies_resolve_entity_metadata_to_breakdown() """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db records = [ _daily_user_spend_record(user_id="user-with-email", api_key="key-1", spend=7.0), @@ -1088,6 +1105,7 @@ async def test_model_groups_breakdown_keys_by_public_name_with_model_fallback(): """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db records = [ _daily_user_spend_record(user_id="u1", api_key="key-1", spend=7.0, model="gpt-5.2", model_group="gpt-5.2-eu"), @@ -1398,6 +1416,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_rows = [ { @@ -1569,6 +1588,7 @@ async def test_get_daily_activity_aggregated_bounds_api_key_rollups( row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) @@ -1638,6 +1658,7 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) @@ -1666,6 +1687,7 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both def _prisma_with_marker(marker: str | None) -> MagicMock: prisma = MagicMock() prisma.db = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) row = ( @@ -1857,6 +1879,7 @@ async def test_get_daily_activity_aggregated_reports_exact_limit_key_count_as_co row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) @@ -1906,6 +1929,7 @@ async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_mo mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, []) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) @@ -2560,6 +2584,7 @@ class TestPtuCostAttributionDisabled: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_table = MagicMock() mock_table.count = AsyncMock(return_value=2) mock_table.find_many = AsyncMock( @@ -2598,6 +2623,7 @@ class TestPtuCostAttributionDisabled: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_table = MagicMock() mock_table.count = AsyncMock(return_value=2) mock_table.find_many = AsyncMock( @@ -2727,6 +2753,7 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): query's rollup dispatch.""" mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db base = { "model": None, @@ -2831,6 +2858,7 @@ async def test_get_api_key_metadata_resolves_session_key_via_spend_log_window(): session_digest = hash_token("cli-session-user-42") mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock( diff --git a/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py index 0ecc4f8d7cb..39d4fd06e75 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py +++ b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py @@ -133,6 +133,7 @@ async def test_migrate_requires_aes_gate(salt_key, monkeypatch): def _config_prisma(record): """Build an AsyncMock prisma client whose litellm_config returns `record`.""" client = MagicMock() + client.replica_db = client.db client.db.litellm_config.find_unique = AsyncMock(return_value=record) client.db.litellm_config.update = AsyncMock() return client @@ -218,6 +219,7 @@ async def test_sso_walker_real_run_migrates_and_clears_residual(salt_key, monkey _enable_aes(monkeypatch) record = SimpleNamespace(sso_settings={"client_secret": legacy, "client_id": "id"}) client = MagicMock() + client.replica_db = client.db client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record) client.db.litellm_ssoconfig.update = AsyncMock() @@ -235,6 +237,7 @@ async def test_sso_walker_dry_run_reports_residual_not_migrated(salt_key, monkey _enable_aes(monkeypatch) record = SimpleNamespace(sso_settings={"client_secret": legacy}) client = MagicMock() + client.replica_db = client.db client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record) client.db.litellm_ssoconfig.update = AsyncMock() @@ -254,6 +257,7 @@ async def test_check_reports_residual_legacy(salt_key, monkeypatch): _enable_aes(monkeypatch) client = MagicMock() + client.replica_db = client.db # Net-new walker tables: empty team / token / sso, one legacy vantage field. client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -278,6 +282,7 @@ async def test_check_reports_residual_legacy(salt_key, monkeypatch): async def test_check_reports_zero_after_migration(salt_key, monkeypatch): _enable_aes(monkeypatch) client = MagicMock() + client.replica_db = client.db client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) @@ -314,6 +319,7 @@ async def test_callback_vars_walker_migrates_team_metadata(salt_key, monkeypatch team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) client = MagicMock() + client.replica_db = client.db client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) client.db.litellm_teamtable.update = AsyncMock() @@ -342,6 +348,7 @@ async def test_callback_vars_walker_dry_run_reports_legacy(salt_key, monkeypatch team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) client = MagicMock() + client.replica_db = client.db client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) client.db.litellm_teamtable.update = AsyncMock() @@ -379,6 +386,7 @@ async def test_callback_vars_walker_migrates_callback_settings_shape( team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) client = MagicMock() + client.replica_db = client.db client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) client.db.litellm_teamtable.update = AsyncMock() @@ -414,6 +422,7 @@ async def test_check_reports_callback_var_legacy_with_gate_off(salt_key, monkeyp team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta) client = MagicMock() + client.replica_db = client.db _empty_covered_tables(client) client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -439,6 +448,7 @@ async def test_scan_covered_tables_classifies_legacy_and_v2(salt_key, monkeypatc v2 = encrypt_value_helper("cred-secret") client = MagicMock() + client.replica_db = client.db _empty_covered_tables(client) client.db.litellm_proxymodeltable.find_many = AsyncMock( return_value=[ @@ -489,6 +499,7 @@ async def test_check_classifies_mcp_secret_maps( value: Final = json.dumps(cases[case]) if as_json and case != "null" else cases[case] row: Final = SimpleNamespace(**{column: value}) client: Final = MagicMock() + client.replica_db = client.db _empty_covered_tables(client) client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row]) client.db.litellm_mcpservertable.update = AsyncMock() @@ -526,6 +537,7 @@ async def test_check_counts_covered_table_residual(salt_key, monkeypatch): _enable_aes(monkeypatch) client = MagicMock() + client.replica_db = client.db _empty_covered_tables(client) client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -554,6 +566,7 @@ async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch) row = SimpleNamespace(litellm_params={"api_key": legacy}) client = MagicMock() + client.replica_db = client.db _empty_covered_tables(client) client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[row]) client.db.litellm_config.find_unique = AsyncMock(return_value=None) diff --git a/tests/test_litellm/proxy/management_endpoints/test_gateway_request_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_gateway_request_endpoints.py index 4f4e378bae4..af0ff8e5353 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_gateway_request_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_gateway_request_endpoints.py @@ -77,6 +77,7 @@ def _admin() -> UserAPIKeyAuth: def _prisma_returning(rows: list) -> MagicMock: client = MagicMock() client.db = MagicMock() + client.replica_db = client.db client.db.query_raw = AsyncMock(return_value=rows) return client diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index c663e63414c..62313c14239 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -55,6 +55,7 @@ async def test_ui_view_users_with_null_email(mocker, caplog): """ # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Create mock user data with null email mock_user = mocker.MagicMock() @@ -97,6 +98,7 @@ async def test_ui_view_users_proxy_admin_no_org_filter(mocker): Proxy admin: find_many is called without organization_memberships in where. """ mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def mock_find_many(*args, **kwargs): assert "organization_memberships" not in (kwargs.get("where") or {}) @@ -165,6 +167,7 @@ def test_ui_view_users_search_matches_user_id_or_email( return [user for user in users if _matches_user_where(user, where)] mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_many = mock_find_many mocker.patch( # test-quality-ok: endpoint reads settings via module global; same seam as sibling tests "litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints.get_ui_settings_cached", @@ -194,6 +197,7 @@ async def test_ui_view_users_org_admin_filtered_by_org(mocker): from litellm.proxy._types import LiteLLM_OrganizationMembershipTable mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db org_id = "org-123" async def mock_find_many(*args, **kwargs): @@ -253,6 +257,7 @@ async def test_ui_view_users_non_org_admin_returns_403(mocker): from fastapi import HTTPException mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Flag ON mocker.patch( @@ -296,6 +301,7 @@ async def test_ui_view_users_flag_off_internal_user_can_search(mocker): Flag OFF (default): any authenticated user can search all users without org filtering. """ mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def mock_find_many(*args, **kwargs): where = kwargs.get("where") or {} @@ -331,6 +337,7 @@ async def test_ui_view_users_flag_on_team_admin_org_team(mocker): from litellm.proxy._types import LiteLLM_TeamTableCachedObj mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db org_id = "org-456" tid = "team-789" @@ -402,6 +409,7 @@ async def test_ui_view_users_flag_on_team_admin_non_org_team_403(mocker): from litellm.proxy._types import LiteLLM_TeamTableCachedObj mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db tid = "team-no-org" # Flag ON @@ -463,6 +471,7 @@ async def test_ui_view_users_flag_on_team_admin_org_member_no_team_id(mocker): should succeed and filter by the user's org membership. """ mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db org_id = "org-member-org" async def mock_find_many(*args, **kwargs): @@ -521,6 +530,7 @@ async def test_ui_view_users_flag_on_team_admin_not_in_org_resolves_via_key_team from litellm.proxy._types import LiteLLM_TeamTableCachedObj mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db org_id = "org-from-team" tid = "key-team-id" @@ -611,6 +621,7 @@ async def test_get_users_includes_timestamps(mocker): """ # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Create mock user data with timestamps mock_user_data = { @@ -677,6 +688,7 @@ async def test_get_users_redacts_scim_enterprise_metadata(mocker): the rest of the metadata intact, matching the user-info endpoints. """ mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_row = mocker.MagicMock() mock_user_row.user_id = "listed-user" @@ -871,6 +883,7 @@ async def test_new_user_license_over_limit(mocker): # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # 1000 billable users (no SCIM-deactivated rows): the filtered count used # for "deactivated" returns 0, so billable == total == 1000 @@ -954,6 +967,7 @@ async def test_new_user_license_gate_counts_only_billable_users(mocker): def _prisma(total, deactivated): client = mocker.MagicMock() + client.replica_db = client.db async def _count(*args, where=None, **kwargs): return deactivated if where is not None else total @@ -991,6 +1005,7 @@ async def test_new_user_non_admin_cannot_create_admin(mocker): # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Setup the mock count response (under license limit) async def mock_count(*args, **kwargs): @@ -1063,6 +1078,7 @@ async def test_new_user_non_admin_permissions_non_empty_rejected(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import new_user mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def mock_count(*args, **kwargs): return 5 @@ -1106,6 +1122,7 @@ async def test_new_user_non_admin_permissions_explicit_empty_rejected(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import new_user mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def mock_count(*args, **kwargs): return 5 @@ -1150,6 +1167,7 @@ async def test_new_user_non_admin_omits_permissions_succeeds(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import new_user mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def mock_count(*args, **kwargs): return 5 @@ -1200,6 +1218,7 @@ async def test_new_user_admin_can_set_permissions(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import new_user mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def mock_count(*args, **kwargs): return 5 @@ -1254,6 +1273,7 @@ async def test_update_single_user_non_admin_permissions_rejected(mocker): ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") @@ -1281,6 +1301,7 @@ async def test_update_single_user_non_admin_permissions_explicit_empty_rejected( ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mocker.patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin") @@ -1310,6 +1331,7 @@ async def test_user_info_url_encoding_plus_character(mocker): # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Create a real LiteLLM_UserTable instance (BaseModel) so isinstance check passes mock_user = LiteLLM_UserTable( @@ -1382,6 +1404,7 @@ async def test_user_info_nonexistent_user(mocker): # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock get_data to return None (user doesn't exist) async def mock_get_data(*args, **kwargs): @@ -1428,6 +1451,7 @@ async def test_user_info_no_user_id_view_only_admin_gets_proxy_admin_payload(moc from litellm.proxy.management_endpoints.internal_user_endpoints import user_info mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.get_data = mocker.AsyncMock(return_value=None) mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -1460,6 +1484,7 @@ async def test_new_user_default_teams_flow(mocker): # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Setup the mock count response (under license limit) async def mock_count(*args, **kwargs): @@ -1705,6 +1730,7 @@ async def test_check_duplicate_user_email_case_insensitive(mocker): # Mock the prisma client mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Test Case 1: Duplicate found with different case # Mock existing user with uppercase email @@ -1770,6 +1796,7 @@ async def test_check_duplicate_user_id(mocker): ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Duplicate user_id should raise mock_existing_user = mocker.MagicMock() @@ -1937,6 +1964,7 @@ async def test_get_users_user_id_partial_match(mocker): from litellm.proxy._types import UserAPIKeyAuth mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_data = { "user_id": "test-user-partial-match", @@ -2034,6 +2062,7 @@ def test_get_users_search_matches_user_id_or_email(mocker): return 0 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_many = mock_find_many mock_prisma_client.db.litellm_usertable.count = mock_count mock_prisma_client.db.litellm_verificationtoken.count = mock_key_count @@ -2186,6 +2215,7 @@ async def test_user_model_budget_update_by_email_refreshes_cached_user(mocker: M max_budget=50, ) prisma_client: Final = mocker.MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user) prisma_client.get_data = mocker.AsyncMock(return_value=[saved_user]) prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": saved_user}) @@ -2225,6 +2255,7 @@ async def test_user_status_update_refreshes_cached_user( metadata={"scim_active": False if active is None else not active, "department": "engineering"}, ) prisma_client: Final = mocker.MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user) prisma_client.get_data = mocker.AsyncMock(return_value=[saved_user]) prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": saved_user}) @@ -2263,6 +2294,7 @@ async def test_bulk_user_model_budget_clear_serializes_and_refreshes_cache(mocke saved_user: Final = LiteLLM_UserTable(user_id="user-spruce", model_max_budget={"model-spruce": {"budget_limit": 5}}) prisma_client: Final = mocker.MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_many = mocker.AsyncMock(return_value=[saved_user]) prisma_client.db.litellm_usertable.update_many = mocker.AsyncMock(return_value=1) mocker.patch("litellm.proxy.proxy_server.prisma_client", prisma_client) # test-quality-ok: substitute the database dependency @@ -2308,6 +2340,7 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker: saved_user: Final = LiteLLM_UserTable(user_id="user-spruce", max_budget=500.0) prisma_client: Final = mocker.MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user) prisma_client.db.litellm_usertable.find_many = mocker.AsyncMock(return_value=[saved_user]) prisma_client.db.litellm_usertable.update_many = mocker.AsyncMock(return_value=1) @@ -2377,6 +2410,7 @@ async def test_get_user_daily_activity_non_admin_cannot_view_other_users(monkeyp # Mock the prisma client so the DB-not-connected check passes mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Non-admin caller @@ -2449,6 +2483,7 @@ async def test_get_user_daily_activity_rejects_service_account_caller(monkeypatc ) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Tripwire: ensure get_daily_activity is never reached @@ -2499,6 +2534,7 @@ async def test_get_user_daily_activity_aggregated_rejects_service_account_caller ) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_get_daily_agg = AsyncMock() @@ -2544,6 +2580,7 @@ async def test_get_user_daily_activity_aggregated_admin_global_view(monkeypatch, # Mock the prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Mock the downstream helper so we don't need a real DB @@ -2610,6 +2647,7 @@ async def test_get_user_daily_activity_aggregated_non_admin_cannot_view_other_us ) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) non_admin_key_dict = UserAPIKeyAuth( @@ -2671,6 +2709,7 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock user lookup mock_user_row = mocker.MagicMock() @@ -2784,6 +2823,7 @@ async def test_delete_user_evicts_jwt_key_mapping_cache_of_its_keys(mocker): user_row.model_dump.return_value = {"user_id": "jwt-user", "user_email": "jwt-user@example.com", "teams": []} mock_prisma_client: Final = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=user_row) mock_prisma_client.db.litellm_teamtable.find_many = mocker.AsyncMock(return_value=[]) mock_prisma_client.db.litellm_jwtkeymapping = jwt_table @@ -2835,6 +2875,7 @@ async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Target user exists and is a member of org-B only. mock_target_user = mocker.MagicMock() @@ -2904,6 +2945,7 @@ async def test_user_update_rejects_silent_create_for_non_proxy_admin(mocker): ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # user_email lookup yields None → would silently create pre-fix. mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=None) mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -2939,6 +2981,7 @@ async def test_user_info_v2_proxy_admin_can_query_any_user(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_row = mocker.MagicMock() mock_user_row.model_dump.return_value = { @@ -3002,6 +3045,7 @@ async def test_user_info_v2_redacts_scim_enterprise_metadata(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_row = mocker.MagicMock() mock_user_row.model_dump.return_value = { @@ -3083,6 +3127,7 @@ async def test_user_info_v2_internal_user_can_query_self(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_row = mocker.MagicMock() mock_user_row.model_dump.return_value = { @@ -3137,6 +3182,7 @@ async def test_user_info_v2_internal_user_cannot_query_other(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Caller user has no teams (so no team admin access) mock_caller_row = mocker.MagicMock() @@ -3177,6 +3223,7 @@ async def test_user_info_v2_no_user_id_defaults_to_self(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_row = mocker.MagicMock() mock_user_row.model_dump.return_value = { @@ -3231,6 +3278,7 @@ async def test_user_info_v2_nonexistent_user_returns_404(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def mock_find_unique(*args, **kwargs): return None @@ -3266,6 +3314,7 @@ async def test_user_info_v2_response_shape(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_row = mocker.MagicMock() mock_user_row.model_dump.return_value = { @@ -3356,6 +3405,7 @@ async def test_user_info_v2_team_admin_can_query_team_member(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Caller (team admin) mock_caller = mocker.MagicMock() @@ -3435,6 +3485,7 @@ async def test_user_info_v2_team_admin_cannot_query_non_team_member(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Caller (team admin of team-A) mock_caller = mocker.MagicMock() @@ -3497,6 +3548,7 @@ async def test_user_info_v2_url_encoding_plus_character(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2 mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db expected_user_id = "machine-user+admin@example.com" @@ -3729,6 +3781,7 @@ async def test_ghsa_wvg4_non_admin_cannot_self_escalate_max_budget(mocker, budge ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": "user-1", "data": {"user_id": "user-1"}}) existing_user = mocker.MagicMock() existing_user.model_dump.return_value = { @@ -3762,6 +3815,7 @@ async def test_ghsa_wvg4_non_admin_cannot_self_escalate_spend(mocker): ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_user = mocker.MagicMock() existing_user.model_dump.return_value = { "user_id": "user-1", @@ -3794,6 +3848,7 @@ async def test_ghsa_wvg4_proxy_admin_can_update_user_budget(mocker): ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_user = mocker.MagicMock() existing_user.model_dump.return_value = { "user_id": "target-user", @@ -3828,6 +3883,7 @@ async def test_admin_user_update_spend_invalidates_counter(mocker): ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_user = mocker.MagicMock() existing_user.model_dump.return_value = {"user_id": "target-user", "spend": 50.0} existing_user.user_id = "target-user" @@ -3863,6 +3919,7 @@ async def test_user_update_rejects_non_finite_spend(mocker): ) mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_user = mocker.MagicMock() existing_user.model_dump.return_value = {"user_id": "target-user", "spend": 50.0} existing_user.user_id = "target-user" @@ -3897,6 +3954,7 @@ async def test_resolve_user_email_metadata_maps_page_user_ids_to_email(mocker): """ mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db find_many = mocker.AsyncMock( return_value=[ SimpleNamespace(user_id="u1", user_email="alice@example.com", user_alias="Alice"), @@ -3925,6 +3983,7 @@ async def test_resolve_user_email_metadata_maps_page_user_ids_to_email(mocker): async def test_resolve_user_email_metadata_skips_db_when_no_user_ids(mocker): """No user_ids on the page (e.g. all spend is unattributed) means no query.""" mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db find_many = mocker.AsyncMock(return_value=[]) mock_prisma_client.db.litellm_usertable.find_many = find_many @@ -4075,6 +4134,7 @@ async def test_get_user_info_for_proxy_admin_validates_keys_and_teams(): ] mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.query_raw = AsyncMock(return_value=raw_rows) with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): @@ -4091,6 +4151,7 @@ async def test_get_user_info_for_proxy_admin_validates_keys_and_teams(): def _object_permission_mocks(mocker, existing_object_permission_id=None): """Prisma double whose user row optionally already links a permission row.""" mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_user = mocker.MagicMock() existing_user.model_dump.return_value = { "user_id": "target-user", @@ -4316,6 +4377,7 @@ async def test_new_user_persists_the_requested_mcp_entitlement(mocker): """generate_key_helper_fn only forwards object_permission_id, so /user/new has to create the grants row itself; otherwise the entitlement the admin sent is silently dropped.""" mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable.create = mocker.AsyncMock( return_value=SimpleNamespace(object_permission_id="perm-created") ) @@ -4459,6 +4521,7 @@ def _admin_prisma(mocker): this file repeats per-test; consolidated here since these three share it verbatim).""" mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mocker.patch( # test-quality-ok: same module-global mocking every test in this file already uses "litellm.proxy.proxy_server.prisma_client", mock_prisma_client ) @@ -4697,6 +4760,7 @@ async def test_delete_user_evicts_cached_user_rows(mocker: MockerFixture) -> Non deleted: Final = LiteLLM_UserTable(user_id="user-gone", user_email="gone@example.test", teams=[]) survivor: Final = LiteLLM_UserTable(user_id="user-stays", user_email="stays@example.test", teams=[]) prisma_client: Final = mocker.MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=deleted) prisma_client.db.litellm_teamtable.find_many = mocker.AsyncMock(return_value=[]) prisma_client.db.litellm_jwtkeymapping.find_many = mocker.AsyncMock(return_value=[]) @@ -4739,6 +4803,7 @@ _DB_OUTAGE_503_BODY: Final = { def _user_read_raising(mocker: MockerFixture, error: Exception) -> tuple[MagicMock, MagicMock]: prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=error) cache = MagicMock() cache.async_get_cache = AsyncMock(return_value=None) @@ -4859,6 +4924,7 @@ async def test_delete_user_writes_deleted_audit_log_for_user_keys(mocker): from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user mock_prisma_client = mocker.MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_row = mocker.MagicMock() mock_user_row.user_id = "doomed-user" diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 3b86f1f6d20..2ba11c8d583 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -92,6 +92,7 @@ client = TestClient(app) @pytest.mark.asyncio async def test_list_keys(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many args = { @@ -128,6 +129,7 @@ async def test_list_keys_include_created_by_keys(): and applies specific filtering to both user's own keys and created_by keys. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_find_many = AsyncMock(return_value=[]) mock_count = AsyncMock(return_value=0) mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many @@ -287,6 +289,7 @@ async def test_key_token_handling(monkeypatch): 2. if token_id exists, it should equal token field """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock( token="hashed_token_123", litellm_budget_table=None, object_permission=None @@ -294,6 +297,7 @@ async def test_key_token_handling(monkeypatch): ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -342,11 +346,13 @@ async def test_budget_reset_and_expires_at_first_of_month(monkeypatch): - expires is set to approximately 1 month from creation time (exact duration) """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None) ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -424,11 +430,13 @@ async def test_key_expiration_exact_duration_hours(monkeypatch): Specifically tests the bug where "12h" duration would expire at midnight instead of 12 hours from creation. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None) ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -487,10 +495,12 @@ async def test_key_expiration_exact_duration_hours(monkeypatch): @pytest.mark.asyncio async def test_generate_key_persists_tpd_limit(monkeypatch): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.insert_data = AsyncMock( return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None) ) mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) @@ -514,6 +524,7 @@ async def test_key_generation_with_object_permission(monkeypatch): """ # --- Setup mocked prisma client --- mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # identity helper for jsonify_object (used inside generate_key_helper_fn) mock_prisma_client.jsonify_object = lambda data: data # type: ignore @@ -523,6 +534,7 @@ async def test_key_generation_with_object_permission(monkeypatch): return_value=MagicMock(object_permission_id="objperm123") ) mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.create = ( mock_object_permission_create @@ -595,8 +607,10 @@ async def test_generate_key_debug_log_never_contains_raw_token(monkeypatch, capl import logging mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db async def _insert_data_side_effect(*args, **kwargs): if kwargs.get("table_name") == "user": @@ -669,8 +683,10 @@ async def test_generate_key_personal_non_admin_denied_for_team_scoped_fields( enforce_member_can_assign_access_groups call in _personal_key_generation_check must break this test.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data # type: ignore mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=MagicMock(object_permission_id="should-not-create") @@ -725,8 +741,10 @@ async def test_update_key_personal_non_admin_denied_vector_stores(monkeypatch): object_permission fields; this test exercises _validate_update_key_data which calls _validate_mcp_servers_for_key_update.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data # type: ignore mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr( "litellm.proxy.proxy_server.user_api_key_cache", @@ -795,6 +813,7 @@ async def test_update_key_grandfathers_existing_mcp_servers(monkeypatch): existing_row.mcp_servers = ["server-a", "server-b"] existing_row.mcp_tool_permissions = {} mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) team_obj = MagicMock() @@ -845,6 +864,7 @@ async def test_update_key_personal_non_admin_denied_access_groups( non-admins. Reverting the enforce move (putting it back inside `if _team_id_to_check is not None`) breaks this test.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data # type: ignore monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -890,8 +910,10 @@ async def test_update_key_personal_non_admin_denied_access_groups( async def test_generate_key_helper_fn_with_access_group_ids(monkeypatch): """Ensure generate_key_helper_fn passes access_group_ids into the key insert payload.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data # type: ignore mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=MagicMock(object_permission_id=None) @@ -941,8 +963,10 @@ async def test_generate_key_helper_fn_with_budget_fallbacks(monkeypatch): kwargs) raised "unexpected keyword argument" before ever reaching the DB. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data # type: ignore mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=MagicMock(object_permission_id=None) @@ -995,6 +1019,7 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch): 3. The key is correctly linked to the object_permission record """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data # Track what data is passed to create @@ -1005,6 +1030,7 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch): return MagicMock(object_permission_id="objperm_mcp_123") mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.create = mock_create mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) @@ -1089,6 +1115,7 @@ async def test_key_update_object_permissions_existing_permission(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock existing key with object_permission_id existing_key_row = LiteLLM_VerificationToken( @@ -1164,6 +1191,7 @@ async def test_key_update_object_permissions_no_existing_permission(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_key_row_no_perm = LiteLLM_VerificationToken( token="test_token_hash_2", @@ -1226,6 +1254,7 @@ async def test_key_update_object_permissions_missing_permission_record(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_key_row_missing_perm = LiteLLM_VerificationToken( token="test_token_hash_3", @@ -1351,6 +1380,7 @@ async def test_key_info_returns_object_permission(monkeypatch): # Mock prisma client mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Mock key with object_permission_id @@ -1446,6 +1476,7 @@ async def test_key_info_returns_lifetime_total_spend_next_to_resettable_spend(mo from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=_stored_key_with_lifetime_spend(token="hashed_key", spend=0.0, total_spend=3.75) @@ -1463,6 +1494,7 @@ async def test_key_info_returns_lifetime_total_spend_next_to_resettable_spend(mo @pytest.mark.asyncio async def test_list_keys_full_object_returns_lifetime_total_spend(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[_stored_key_with_lifetime_spend(token="hashed_key", spend=0.0, total_spend=3.75)] ) @@ -1557,6 +1589,7 @@ async def test_generate_key_fn_rejects_short_custom_key(monkeypatch, short_key): accepted and fully exposed via key_name.""" mock_prisma_client = AsyncMock() mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -1588,11 +1621,13 @@ async def test_generate_key_fn_rejects_short_custom_key(monkeypatch, short_key): async def test_generate_key_fn_accepts_custom_key_at_minimum_length(monkeypatch): """Custom keys at exactly the minimum length (16 chars) are still accepted.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None, object_permission=None) ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -1850,6 +1885,7 @@ async def test_generate_service_account_works_with_team_id(): "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn" ) as mock_generate_key, ): + mock_prisma.replica_db = mock_prisma.db # Configure mocks mock_prisma.return_value = AsyncMock() @@ -1884,6 +1920,7 @@ async def test_generate_key_throttle_rejected_for_non_admin(): /key/update gate does not cover generate, so generate needs its own admin check. Only the enable value is gated, so this must 403.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): with pytest.raises(HTTPException) as exc: await _common_key_generation_helper( @@ -1932,6 +1969,7 @@ async def test_generate_key_end_user_budget_id_rejected_for_non_admin(): """A key's default end-user budget overrides the proxy-wide one, so a non-admin must not be able to pick a looser one for the customers their key creates.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock() with pytest.raises(HTTPException) as exc: await _validate_end_user_budget_id_change( @@ -1965,6 +2003,7 @@ async def test_generate_key_end_user_budget_id_must_name_an_existing_budget(): """A typo in end_user_budget_id would silently leave new customers on the proxy-wide default, so key creation rejects an id that matches no budget row.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) with pytest.raises(HTTPException) as exc: await _validate_end_user_budget_id_change( @@ -1986,6 +2025,7 @@ async def test_generate_key_end_user_budget_id_lands_in_key_metadata(): budget_row = MagicMock() budget_row.model_dump.return_value = {"budget_id": "svc-a-budget", "max_budget": 0.5} mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row) with ( patch( # test-quality-ok: the helper reads proxy_server globals, no seam @@ -2037,6 +2077,7 @@ async def test_update_key_clears_end_user_budget_id_with_empty_string(): existing_key = LiteLLM_VerificationToken(token="hashed", metadata={"end_user_budget_id": "svc-a-budget"}) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) await _validate_update_key_data( @@ -2063,6 +2104,7 @@ async def test_update_key_metadata_body_without_end_user_budget_id_is_a_clear_fo would detach the key default; that must be refused like an explicit clear, while an admin may do it.""" existing_key = LiteLLM_VerificationToken(token="hashed", metadata={"end_user_budget_id": "svc-a-budget"}) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) non_admin = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-alice", user_id="alice") @@ -2103,6 +2145,7 @@ async def test_regenerate_key_end_user_budget_id_rejected_for_non_admin(): from litellm.proxy._types import RegenerateKeyRequest mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock() with pytest.raises(HTTPException) as exc: await _execute_virtual_key_regeneration( @@ -2501,6 +2544,7 @@ async def test_validate_team_id_used_in_service_account_request_requires_team_id ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Test that HTTPException is raised when team_id is None with pytest.raises(HTTPException) as exc_info: @@ -2547,6 +2591,7 @@ async def test_validate_team_id_used_in_service_account_request_checks_team_exis ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock the database query to return None (team doesn't exist) mock_find_unique = AsyncMock(return_value=None) @@ -2577,6 +2622,7 @@ async def test_validate_team_id_used_in_service_account_request_success(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock the database query to return a team object (team exists) mock_team = {"team_id": "existing-team-id", "team_name": "Test Team"} @@ -2609,6 +2655,7 @@ async def test_generate_service_account_key_endpoint_validation(): # Test case 1: Missing team_id with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db # Mock prisma_client to be not None so we can reach team_id validation mock_prisma_instance = AsyncMock() mock_prisma.return_value = mock_prisma_instance @@ -2629,6 +2676,7 @@ async def test_generate_service_account_key_endpoint_validation(): # Test case 2: Team doesn't exist in database with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db # Mock team not found mock_find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_teamtable.find_unique = mock_find_unique @@ -2660,6 +2708,7 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch): # Mock dependencies mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -2768,6 +2817,7 @@ async def test_unblock_key_invalid_key_format(monkeypatch): # Mock prisma_client to avoid DB connection error mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Mock request and user auth @@ -2806,6 +2856,7 @@ async def test_block_key_nonexistent_key_returns_404(monkeypatch): from litellm.proxy.management_endpoints.key_management_endpoints import block_key mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -2862,6 +2913,7 @@ async def test_unblock_key_nonexistent_key_returns_404(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -2916,6 +2968,7 @@ async def test_update_key_nonexistent_key_returns_404(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -2978,6 +3031,7 @@ async def test_update_key_rejects_a_duration_that_never_advances(monkeypatch, ba key_in_db = LiteLLM_VerificationToken(token=hashed_token, user_id="test-user") mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=key_in_db ) @@ -3008,6 +3062,7 @@ async def test_generate_key_rejects_a_duration_that_never_advances(monkeypatch, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) @@ -3048,6 +3103,7 @@ async def test_update_key_by_alias_only(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[key_in_db] ) @@ -3100,6 +3156,7 @@ async def test_update_key_changed_alias_must_match_key_alias_pattern(monkeypatch key_in_db = LiteLLM_VerificationToken(token=hashed_token, key_alias="Legacy Alias", user_id="test-user") mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_in_db) mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[key_in_db]) mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(return_value=None) @@ -3143,6 +3200,7 @@ async def test_update_key_by_alias_not_found_returns_404(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[] ) @@ -3180,6 +3238,7 @@ async def test_update_key_by_duplicate_alias_returns_400(monkeypatch): LiteLLM_VerificationToken(token="hashed-token-2", key_alias="dup-alias"), ] mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=rows ) @@ -3221,6 +3280,7 @@ async def test_update_key_with_key_and_alias_selects_by_key(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=key_in_db ) @@ -3263,6 +3323,7 @@ async def test_block_key_existing_key_succeeds(monkeypatch): from litellm.proxy.management_endpoints.key_management_endpoints import block_key mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -3585,6 +3646,7 @@ async def test_check_team_key_limits_no_existing_keys(): """ # Mock prisma client mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[] ) @@ -3643,6 +3705,7 @@ async def test_check_team_key_limits_with_existing_keys_within_bounds(): existing_key3.rpm_limit = None # Should be ignored in calculation mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key1, existing_key2, existing_key3] ) @@ -3692,6 +3755,7 @@ async def test_check_team_key_limits_tpm_overallocation(): existing_key2.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key1, existing_key2] ) @@ -3749,6 +3813,7 @@ async def test_check_team_key_limits_rpm_overallocation(): existing_key2.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key1, existing_key2] ) @@ -3813,6 +3878,7 @@ async def test_check_team_key_limits_on_update_excludes_self(): other_key.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[self_key, other_key] ) @@ -3859,6 +3925,7 @@ async def test_check_team_key_limits_no_team_limits(): existing_key.rpm_limit = 500 mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key] ) @@ -3902,6 +3969,7 @@ async def test_check_team_key_limits_no_key_limits(): existing_key.rpm_limit = 800 mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key] ) @@ -3955,6 +4023,7 @@ async def test_check_team_key_limits_mixed_scenarios(): existing_key3.rpm_limit = 300 mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key1, existing_key2, existing_key3] ) @@ -3998,6 +4067,7 @@ async def test_check_team_key_limits_exact_boundary(): existing_key.rpm_limit = 700 mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key] ) @@ -4317,6 +4387,7 @@ async def test_generate_key_with_object_permission(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda x: x # Mock object permission creation @@ -4652,6 +4723,7 @@ async def test_check_org_key_limits_no_existing_keys(): """ # Mock prisma client mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[] ) @@ -4715,6 +4787,7 @@ async def test_check_org_key_limits_with_existing_keys_within_bounds(): existing_key3.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key1, existing_key2, existing_key3] ) @@ -4768,6 +4841,7 @@ async def test_check_org_key_limits_tpm_overallocation(): existing_key2.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key1, existing_key2] ) @@ -4827,6 +4901,7 @@ async def test_check_org_key_limits_rpm_overallocation(): existing_key2.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key1, existing_key2] ) @@ -4881,6 +4956,7 @@ async def test_check_org_key_limits_no_org_limits(): existing_key.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key] ) @@ -5310,6 +5386,7 @@ def test_transform_verification_tokens_to_deleted_records_empty_list(): @pytest.mark.asyncio async def test_save_deleted_verification_token_records(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_create_many = AsyncMock() mock_prisma_client.db.litellm_deletedverificationtoken.create_many = ( mock_create_many @@ -5340,6 +5417,7 @@ async def test_save_deleted_verification_token_records(): @pytest.mark.asyncio async def test_save_deleted_verification_token_records_empty_list(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_create_many = AsyncMock() mock_prisma_client.db.litellm_deletedverificationtoken.create_many = ( mock_create_many @@ -5355,6 +5433,7 @@ async def test_save_deleted_verification_token_records_empty_list(): @pytest.mark.asyncio async def test_persist_deleted_verification_tokens(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_create_many = AsyncMock() mock_prisma_client.db.litellm_deletedverificationtoken.create_many = ( mock_create_many @@ -5404,6 +5483,7 @@ async def test_persist_deleted_verification_tokens(): @pytest.mark.asyncio async def test_delete_verification_tokens_persists_deleted_keys(monkeypatch): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() user_api_key_dict = UserAPIKeyAuth( @@ -5567,6 +5647,7 @@ async def test_delete_verification_tokens_evicts_jwt_key_mapping_cache(monkeypat ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[key1] ) @@ -5618,6 +5699,7 @@ async def test_delete_key_fn_persists_deleted_keys(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() user_api_key_dict = UserAPIKeyAuth( @@ -5706,6 +5788,7 @@ async def test_can_delete_verification_token_proxy_admin_team_key(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() async def mock_get_team_object(*args, **kwargs): @@ -5757,6 +5840,7 @@ async def test_can_delete_verification_token_team_admin_different_team(monkeypat ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() async def mock_get_team_object(*args, **kwargs): @@ -5807,6 +5891,7 @@ async def test_can_delete_verification_token_key_owner_team_key(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() async def mock_get_team_object(*args, **kwargs): @@ -5843,6 +5928,7 @@ async def test_can_delete_verification_token_key_owner_personal_key(monkeypatch) ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() result = await can_modify_verification_token( @@ -5887,6 +5973,7 @@ async def test_can_delete_verification_token_other_user_team_key(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() async def mock_get_team_object(*args, **kwargs): @@ -5923,6 +6010,7 @@ async def test_can_delete_verification_token_other_user_personal_key(monkeypatch ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() result = await can_modify_verification_token( @@ -5951,6 +6039,7 @@ async def test_can_delete_verification_token_team_key_no_team_found(monkeypatch) ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() async def mock_get_team_object(*args, **kwargs): @@ -5987,6 +6076,7 @@ async def test_can_delete_verification_token_personal_key_no_user_id(monkeypatch ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() result = await can_modify_verification_token( @@ -6015,6 +6105,7 @@ async def test_can_modify_verification_token_proxy_admin_team_key(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() result = await can_modify_verification_token( @@ -6043,6 +6134,7 @@ async def test_can_modify_verification_token_proxy_admin_personal_key(monkeypatc ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() result = await can_modify_verification_token( @@ -6061,6 +6153,7 @@ async def test_list_keys_with_expand_user(): Test that expand=user parameter correctly includes user information in the response. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Create mock keys with user_ids key1_dict = { @@ -6201,6 +6294,7 @@ async def test_list_keys_with_expand_user_includes_created_by_user(): Test that expand=user also resolves created_by to a user object. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Key created by user789 but owned by user123 key1_dict = { @@ -6295,6 +6389,7 @@ async def test_list_keys_with_status_deleted(): Test that status="deleted" parameter correctly queries the deleted keys table. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock deleted keys table mock_deleted_key1 = MagicMock() @@ -6372,6 +6467,7 @@ async def test_list_keys_with_invalid_status(): from unittest.mock import Mock, patch mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock the endpoint function directly to test validation from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth @@ -6407,6 +6503,7 @@ async def test_list_keys_accepts_live_status_filters(monkeypatch, status_filter) live_row = MagicMock() live_row.model_dump.return_value = {"token": "hashed_live_token", "object_permission_id": None} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[live_row]) mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=1) mock_prisma_client.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) @@ -6481,6 +6578,7 @@ def test_build_key_filter_conditions_deleted_status_adds_no_live_clause(): @pytest.mark.asyncio async def test_list_key_helper_revoked_status_filters_live_table_on_blocked(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) @@ -6529,6 +6627,7 @@ async def test_info_key_fn_serves_deleted_key_from_archive(monkeypatch): hashed = "hashed_deleted_token" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_deletedverificationtoken.find_first = AsyncMock( @@ -6559,6 +6658,7 @@ async def test_info_key_fn_archived_key_keeps_owner_authorization(monkeypatch): hashed = "hashed_deleted_token" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_deletedverificationtoken.find_first = AsyncMock( @@ -6586,6 +6686,7 @@ async def test_info_key_fn_unknown_key_still_404s(monkeypatch): from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_deletedverificationtoken.find_first = AsyncMock(return_value=None) @@ -6614,6 +6715,7 @@ async def test_info_key_fn_reports_live_key_status(monkeypatch, blocked, expires from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) live_row = MagicMock(spec=LiteLLM_VerificationToken) live_row.model_dump.return_value = { @@ -6644,6 +6746,7 @@ async def test_list_keys_non_admin_user_id_auto_set(): from unittest.mock import Mock, patch mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Create a non-admin user with a user_id test_user_id = "test-user-123" @@ -6735,6 +6838,7 @@ async def _invoke_list_keys_and_capture_helper_kwargs( from litellm.proxy._types import LiteLLM_UserTable mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_info = LiteLLM_UserTable( user_id=user_api_key_dict.user_id, user_email="member@example.com", @@ -7177,6 +7281,7 @@ def test_build_key_filter_conditions_search_narrows_team_admin_visibility(): async def test_list_key_helper_applies_search_to_prisma_where(): """LIT-4741: `search` given to _list_key_helper must reach the Prisma where clause.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0) @@ -7210,6 +7315,7 @@ async def _run_bulk_update_on_one_key( token=_BULK_UPDATE_TOKEN, user_id="test-user", team_id="team-1", max_budget=100.0, budget_id="budget-1" ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_in_db) mock_prisma_client.get_data = AsyncMock(return_value=key_in_db) mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) @@ -7361,6 +7467,7 @@ async def test_generate_key_with_router_settings(monkeypatch): 3. Storing router_settings in the key record """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data # Mock prisma_client.insert_data for both user and key tables @@ -7378,6 +7485,7 @@ async def test_generate_key_with_router_settings(monkeypatch): mock_prisma_client.insert_data = AsyncMock(side_effect=_insert_data_side_effect) mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -7483,7 +7591,7 @@ async def test_update_key_with_router_settings( model = SimpleNamespace(model_id="weighted-id", model_name="gpt-4", model_info={}) table = SimpleNamespace(find_many=AsyncMock(return_value=[model])) - db = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table)) + db = SimpleNamespace(db=SimpleNamespace(litellm_proxymodeltable=table), replica_db=SimpleNamespace(litellm_proxymodeltable=table)) # Mock existing key existing_key = LiteLLM_VerificationToken( @@ -7597,6 +7705,7 @@ async def test_get_and_validate_existing_key(): # Test Case 1: Successfully retrieve existing key mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_key = LiteLLM_VerificationToken( token="test-key-123", user_id="user-123", @@ -7663,6 +7772,7 @@ async def test_process_single_key_update(): # Setup mocks mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() mock_llm_router = MagicMock() @@ -7787,6 +7897,7 @@ async def test_bulk_update_keys_success(monkeypatch): # Setup mocks mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() mock_llm_router = MagicMock() @@ -7933,6 +8044,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch): # Setup mocks mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() mock_llm_router = MagicMock() @@ -8150,6 +8262,7 @@ def test_validate_reset_spend_value_none_spend(): @pytest.mark.asyncio async def test_reset_key_spend_success(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -8256,6 +8369,7 @@ async def test_reset_key_spend_resets_budget_windows(monkeypatch): every request even though the key's own reported spend read $0. """ mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -8372,6 +8486,7 @@ async def test_reset_key_spend_no_budget_limits_skips_window_reset(monkeypatch): """A key with no budget_limits must not trigger any extra DB write beyond the lifetime spend update; _reset_key_budget_windows should be a no-op.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -8487,6 +8602,7 @@ async def test_update_key_spend_updates_counter(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = AsyncMock() mock_proxy_logging_obj = MagicMock() @@ -8557,6 +8673,7 @@ async def test_update_key_spend_updates_counter(monkeypatch): async def test_reset_key_spend_success_team_admin(monkeypatch): """Test that team admin can reset key spend for keys in their team.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -8651,6 +8768,7 @@ async def test_reset_key_spend_success_team_admin(monkeypatch): @pytest.mark.asyncio async def test_reset_key_spend_key_not_found(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None ) @@ -8705,6 +8823,7 @@ async def test_reset_key_spend_db_not_connected(monkeypatch): @pytest.mark.asyncio async def test_reset_key_spend_validation_error(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db key_in_db = LiteLLM_VerificationToken( token="hashed-key", user_id="test-user", @@ -8743,6 +8862,7 @@ async def test_reset_key_spend_validation_error(monkeypatch): @pytest.mark.asyncio async def test_reset_key_spend_authorization_failure(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() hashed_key = "hashed-test-key" @@ -8795,6 +8915,7 @@ async def test_reset_key_spend_authorization_failure(monkeypatch): @pytest.mark.asyncio async def test_reset_key_spend_hashed_key(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -8863,6 +8984,7 @@ async def test_reset_key_spend_hashed_key(monkeypatch): @pytest.mark.asyncio async def test_validate_key_list_check_proxy_admin(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db user_api_key_dict = UserAPIKeyAuth( user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user", @@ -8884,6 +9006,7 @@ async def test_validate_key_list_check_proxy_admin(): @pytest.mark.asyncio async def test_validate_key_list_check_team_admin_success(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db user_info = LiteLLM_UserTable( user_id="test-user", user_email="test@example.com", @@ -8917,6 +9040,7 @@ async def test_validate_key_list_check_team_admin_success(): @pytest.mark.asyncio async def test_validate_key_list_check_team_admin_fail(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db user_info = LiteLLM_UserTable( user_id="test-user", user_email="test@example.com", @@ -8951,6 +9075,7 @@ async def test_validate_key_list_check_team_admin_fail(): @pytest.mark.asyncio async def test_validate_key_list_check_key_hash_authorized(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db user_info = LiteLLM_UserTable( user_id="test-user", user_email="test@example.com", @@ -8997,6 +9122,7 @@ async def test_validate_key_list_check_key_hash_authorized(): @pytest.mark.asyncio async def test_validate_key_list_check_key_hash_unauthorized(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db user_info = LiteLLM_UserTable( user_id="test-user", user_email="test@example.com", @@ -9044,6 +9170,7 @@ async def test_validate_key_list_check_key_hash_unauthorized(): @pytest.mark.asyncio async def test_validate_key_list_check_key_hash_not_found(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db user_info = LiteLLM_UserTable( user_id="test-user", user_email="test@example.com", @@ -9083,6 +9210,7 @@ async def test_validate_key_list_check_key_hash_row_missing(): """A key_hash with no row reaches the same 'Key Hash not found' 403 as a failed lookup, instead of blowing up inside the ownership check on a None row.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=LiteLLM_UserTable( user_id="test-user", @@ -9122,6 +9250,7 @@ async def test_validate_key_list_check_proxy_admin_viewer_skips_db_lookup(): """proxy_admin_viewer takes the same unscoped read fast-path as proxy_admin, so no user row is fetched and none of the user/team scoping filters apply.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=LiteLLM_UserTable( user_id="viewer-user", @@ -9156,6 +9285,7 @@ async def test_validate_key_list_check_internal_user_cannot_query_other_user(): """Admin-view parity must not leak past the admin roles: an internal user still cannot list another user's keys.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=LiteLLM_UserTable( user_id="test-user", @@ -9199,6 +9329,7 @@ async def test_key_with_budget_id_does_not_store_budget_duration(): from unittest.mock import AsyncMock, MagicMock, patch mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_generate_key = AsyncMock( return_value={ @@ -9257,6 +9388,7 @@ async def test_key_does_not_override_explicit_budget_duration(): from unittest.mock import AsyncMock, MagicMock, patch mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # The budget tier has budget_duration="7d" mock_budget_row = MagicMock() mock_budget_row.budget_duration = "7d" @@ -9337,6 +9469,7 @@ async def test_rotate_master_key_reencrypts_model_params_in_place( # Setup mock prisma client mock_prisma_client = AsyncMock() mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock model table — return one model mock_model = MagicMock() @@ -9446,6 +9579,7 @@ async def test_default_key_generate_params_duration(monkeypatch): import litellm mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock( token="hashed_token_123", litellm_budget_table=None, object_permission=None @@ -9453,6 +9587,7 @@ async def test_default_key_generate_params_duration(monkeypatch): ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -9498,6 +9633,7 @@ async def test_default_key_generate_params_object_permission_applied_when_absent import litellm mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock( token="hashed_token_123", litellm_budget_table=None, object_permission=None @@ -9505,6 +9641,7 @@ async def test_default_key_generate_params_object_permission_applied_when_absent ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -9561,6 +9698,7 @@ async def test_default_key_generate_params_object_permission_merges_partial( from litellm.proxy._types import LiteLLM_ObjectPermissionBase mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock( token="hashed_token_123", litellm_budget_table=None, object_permission=None @@ -9568,6 +9706,7 @@ async def test_default_key_generate_params_object_permission_merges_partial( ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -9626,6 +9765,7 @@ async def test_default_key_generate_params_object_permission_does_not_override_e from litellm.proxy._types import LiteLLM_ObjectPermissionBase mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock( token="hashed_token_123", litellm_budget_table=None, object_permission=None @@ -9633,6 +9773,7 @@ async def test_default_key_generate_params_object_permission_does_not_override_e ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -9694,6 +9835,7 @@ async def test_default_key_generate_params_object_permission_not_rejected_for_no import litellm mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert_data = AsyncMock( return_value=MagicMock( token="hashed_token_123", litellm_budget_table=None, object_permission=None @@ -9701,6 +9843,7 @@ async def test_default_key_generate_params_object_permission_not_rejected_for_no ) mock_prisma_client.insert_data = mock_insert_data mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=None @@ -10288,6 +10431,7 @@ async def test_get_member_team_ids(): # Mock prisma client mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Create mock team objects - user is admin of team-A, member of team-B, not in team-C's members list mock_team_a = MagicMock() @@ -10377,6 +10521,7 @@ async def test_generate_key_helper_fn_agent_id(): import litellm.proxy.management_endpoints.key_management_endpoints as km mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_insert = AsyncMock( return_value=MagicMock( token="sk-test", @@ -10416,6 +10561,7 @@ def _make_admin_key_dict() -> UserAPIKeyAuth: async def test_key_aliases_response_shape(): """Test that key_aliases returns the correct paginated response shape.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.query_raw = AsyncMock( side_effect=[ [{"count": 2}], @@ -10448,6 +10594,7 @@ async def test_key_aliases_response_shape(): async def test_key_aliases_pagination_skip_take(): """Test that LIMIT and OFFSET are correctly derived from page and size.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.query_raw = AsyncMock( side_effect=[ [{"count": 120}], @@ -10478,6 +10625,7 @@ async def test_key_aliases_pagination_skip_take(): async def test_key_aliases_search_filter(): """Test that the search param adds a case-insensitive ILIKE condition.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.query_raw = AsyncMock( side_effect=[ [{"count": 0}], @@ -10505,6 +10653,7 @@ async def test_key_aliases_search_filter(): async def test_key_aliases_no_search_omits_ilike_filter(): """Test that without a search term no ILIKE condition is added.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.query_raw = AsyncMock( side_effect=[ [{"count": 0}], @@ -10528,6 +10677,7 @@ async def test_key_aliases_no_search_omits_ilike_filter(): async def test_key_aliases_internal_user_scoped_to_own_keys_and_teams(): """Test that internal users only see aliases for their own keys and team keys.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock user table lookup to return teams mock_user_row = MagicMock() @@ -10574,6 +10724,7 @@ async def test_key_aliases_internal_user_scoped_to_own_keys_and_teams(): async def test_key_aliases_admin_sees_all(): """Test that proxy admins see all aliases without user/team scoping.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.query_raw = AsyncMock( side_effect=[ [{"count": 3}], @@ -10761,6 +10912,7 @@ async def test_check_org_key_limits_on_update_within_bounds(): a key's TPM/RPM limits within organization bounds. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[] ) @@ -10815,6 +10967,7 @@ async def test_check_org_key_limits_on_update_overallocation(): existing_key.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[existing_key] ) @@ -10876,6 +11029,7 @@ async def test_check_org_key_limits_on_update_excludes_self(): other_key.metadata = {} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[self_key, other_key] ) @@ -10977,6 +11131,7 @@ def test_update_key_request_has_organization_id(): def _setup_block_unblock_mocks(monkeypatch, mock_key_team_id=None): """Helper to set up common mocks for block/unblock tests.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = MagicMock() mock_proxy_logging_obj = MagicMock() @@ -11155,6 +11310,7 @@ async def test_update_key_max_budget_rejected_for_internal_user(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = AsyncMock() mock_proxy_logging_obj = MagicMock() @@ -11220,6 +11376,7 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = AsyncMock() mock_proxy_logging_obj = MagicMock() @@ -11323,6 +11480,7 @@ async def test_update_key_throttle_on_budget_exceeded_rejected_for_internal_user ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = AsyncMock() mock_proxy_logging_obj = MagicMock() @@ -11393,6 +11551,7 @@ async def test_update_key_throttle_unchanged_allows_non_budget_edit_for_internal ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_cache = AsyncMock() mock_proxy_logging_obj = MagicMock() @@ -11479,6 +11638,7 @@ async def test_update_key_non_budget_rejects_cross_user_modification(monkeypatch ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db test_hashed_token = "cafebabe" * 8 mock_existing_key = MagicMock() @@ -11542,6 +11702,7 @@ async def test_update_key_creator_reassigned_key_blocked(monkeypatch): test_hashed_token = "aabbccdd" * 8 mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_existing_key = MagicMock() mock_existing_key.token = test_hashed_token @@ -11646,6 +11807,7 @@ async def test_update_key_team_member_with_permission_can_update_non_budget( mock_updated_key.key_alias = "renamed-by-member" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) mock_prisma_client.update_data = AsyncMock(return_value=mock_updated_key) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( @@ -11753,6 +11915,7 @@ async def test_update_key_team_member_cannot_change_budget(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=mock_existing_key @@ -11813,6 +11976,7 @@ class TestLIT1884KeyGenerateValidation: _common_key_generation_helper. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db data = GenerateKeyRequest(key_alias="test-alias") assert data.user_id is None @@ -11850,6 +12014,7 @@ class TestLIT1884KeyGenerateValidation: key/generate should raise ProxyException with status 400. """ mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db data = GenerateKeyRequest( key_alias="test-alias", @@ -11897,6 +12062,7 @@ class TestLIT1884KeyGenerateValidation: ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -11935,6 +12101,7 @@ class TestLIT1884KeyGenerateValidation: ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -12089,6 +12256,7 @@ class TestLIT1884KeyUpdateValidation: ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Should NOT raise await _validate_update_key_data( @@ -12485,6 +12653,7 @@ class TestKeyAliasSkipValidationOnUnchanged: def mock_prisma(self): prisma = MagicMock() prisma.db = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_verificationtoken = MagicMock() prisma.get_data = AsyncMock(return_value=None) # no duplicate alias prisma.update_data = AsyncMock(return_value=None) @@ -12747,6 +12916,7 @@ def _make_regenerate_mock_prisma(): return iter(self._data.items()) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock( return_value=DictLikeResult( { @@ -13311,6 +13481,7 @@ def _policy_existing_team_key() -> LiteLLM_VerificationToken: def _setup_update_key_fn_policy_mocks(monkeypatch, existing_key: LiteLLM_VerificationToken) -> AsyncMock: mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key) mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock(return_value=None) mock_prisma_client.update_data = AsyncMock(return_value={"data": {"max_budget": 50.0, "team_id": "team-a"}}) @@ -13421,6 +13592,7 @@ async def _process_single_key_update_under_policy(prisma_client: AsyncMock, data @pytest.mark.asyncio async def test_process_single_key_update_runs_custom_key_policy_on_the_effective_row(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db updated_row = MagicMock() updated_row.model_dump.return_value = {"max_budget": 50.0, "team_id": "team-a"} mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_row}) @@ -13438,6 +13610,7 @@ async def test_process_single_key_update_runs_custom_key_policy_on_the_effective @pytest.mark.asyncio async def test_process_single_key_update_rejects_when_custom_key_policy_denies(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.update_data = AsyncMock() received: list[CustomKeyPolicyRequest] = [] data = UpdateKeyRequest(key=_POLICY_HASHED_TOKEN, duration="3000d", max_budget=50.0) @@ -13537,6 +13710,7 @@ async def test_update_key_fn_denied_by_the_policy_leaves_the_object_permission_r @pytest.mark.asyncio async def test_process_single_key_update_writes_the_object_permission_row_only_after_the_policy_allows(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db updated_row = MagicMock() updated_row.model_dump.return_value = {"max_budget": 50.0, "team_id": "team-a"} mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_row}) @@ -13553,6 +13727,7 @@ async def test_process_single_key_update_writes_the_object_permission_row_only_a @pytest.mark.asyncio async def test_process_single_key_update_denied_by_the_policy_leaves_the_object_permission_row_untouched(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.update_data = AsyncMock() events: list[str] = [] _record_object_permission_writes(mock_prisma_client, events) @@ -13613,6 +13788,7 @@ async def test_bulk_update_keys_runs_custom_key_policy_per_key(monkeypatch): updated_row = MagicMock() updated_row.model_dump.return_value = {"user_id": "user-123", "max_budget": 100.0} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(side_effect=existing_keys) mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_row}) mock_prisma_client.get_data = AsyncMock(return_value=None) @@ -13667,6 +13843,7 @@ async def test_bulk_update_keys_runs_custom_key_policy_per_key(monkeypatch): def _policy_generate_prisma() -> MagicMock: mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_budgettable.create = AsyncMock(return_value=MagicMock(budget_id="budget-1")) mock_prisma.jsonify_object = MagicMock(side_effect=lambda data: json.loads(data) if isinstance(data, str) else data) return mock_prisma @@ -14239,6 +14416,7 @@ class TestAllowedRoutesCallerPermission: user_role=LitellmUserRoles.INTERNAL_USER, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -14271,6 +14449,7 @@ class TestAllowedRoutesCallerPermission: user_role=LitellmUserRoles.PROXY_ADMIN, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db stub_response = MagicMock() with ( @@ -14303,6 +14482,7 @@ class TestAllowedRoutesCallerPermission: user_role=LitellmUserRoles.INTERNAL_USER, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db stub_response = MagicMock() with ( @@ -14334,6 +14514,7 @@ class TestAllowedRoutesCallerPermission: user_role=LitellmUserRoles.INTERNAL_USER, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -14375,6 +14556,7 @@ class TestAllowedRoutesCallerPermission: user_role=LitellmUserRoles.INTERNAL_USER, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -14414,6 +14596,7 @@ class TestAllowedRoutesCallerPermission: user_role=LitellmUserRoles.INTERNAL_USER, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -14507,6 +14690,7 @@ class TestAllowedRoutesCallerPermission: user_role=LitellmUserRoles.INTERNAL_USER, ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -14669,6 +14853,7 @@ async def test_process_single_key_update_cache_invalidation_with_token_hash(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=existing_key ) @@ -14809,6 +14994,7 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db # _execute_virtual_key_regeneration calls dict(updated_token) which # needs the return value to be iterable as key-value pairs. @@ -14926,6 +15112,7 @@ def _setup_team_keys_mocks( ): """Set up mocks for bulk_update_team_keys; returns mock_prisma.""" mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[] if find_many is None else find_many ) @@ -15679,6 +15866,7 @@ async def test_regenerate_applies_normalized_mcp_object_permission(): ) existing_key = _make_regenerate_existing_key() mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_repo = MagicMock() mock_repo.table.find_unique = AsyncMock(return_value=existing_key) execute_mock = AsyncMock(return_value=MagicMock()) @@ -15758,6 +15946,7 @@ async def test_ghsa_q775_non_admin_unlimited_can_delegate_budget(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -15792,6 +15981,7 @@ async def test_ghsa_q775_non_admin_cannot_exceed_own_budget(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -15825,6 +16015,7 @@ async def test_ghsa_q775_non_admin_within_budget_allowed(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -15861,6 +16052,7 @@ async def test_ghsa_q775_upperbound_default_not_rejected(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -15901,6 +16093,7 @@ async def test_ghsa_q775_default_key_generate_params_not_rejected(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -15938,6 +16131,7 @@ async def test_ghsa_q775_admin_bypasses_budget_ceiling(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -16022,6 +16216,7 @@ async def test_ghsa_q775_ui_session_token_personal_key_still_capped(): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), @@ -16258,6 +16453,7 @@ async def test_info_key_fn_includes_model_max_budget_usage(monkeypatch): } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -16320,6 +16516,7 @@ async def test_info_key_fn_no_model_max_budget_skips_usage(monkeypatch): test_key_token = "hashed_token_no_budget" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() mock_user_api_key_cache.async_get_cache = AsyncMock() @@ -16381,6 +16578,7 @@ async def test_info_key_fn_v2_includes_model_max_budget_usage(monkeypatch): model_max_budget = {"gpt-4o": {"budget_limit": 1.00, "time_period": "7d"}} mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -16445,6 +16643,7 @@ async def test_info_key_fn_budget_table_fallback(monkeypatch): } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -16518,6 +16717,7 @@ async def test_info_key_fn_v2_budget_table_fallback(monkeypatch): } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -16591,6 +16791,7 @@ async def test_info_key_fn_reports_budget_limits_usage(monkeypatch): ] mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -16655,6 +16856,7 @@ async def test_info_key_fn_no_budget_limits_skips_spend_lookup(monkeypatch): test_key_token = "hashed_token_no_windows" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -16725,6 +16927,7 @@ async def test_info_key_fn_v2_reports_budget_limits_usage(monkeypatch): ] mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -16893,6 +17096,7 @@ async def test_info_key_fn_reads_the_configured_budget_model_key(monkeypatch): } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) mock_user_api_key_cache = AsyncMock() monkeypatch.setattr( @@ -17234,6 +17438,7 @@ def _list_team_a_keys_as(user_role, members_with_roles, query): from litellm.proxy.management_endpoints.key_management_endpoints import router mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = _InMemoryVerificationTokenTable(_TEAM_A_KEYS) mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=LiteLLM_UserTable(user_id="alice", teams=["team-a"], organization_memberships=[]) @@ -17807,6 +18012,7 @@ async def test_update_key_non_admin_permissions_non_empty_rejected(monkeypatch): """`_validate_update_key_data` rejects a non-admin when `permissions` is present in the request body (personal-key fast-path caller).""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -17835,6 +18041,7 @@ async def test_update_key_non_admin_permissions_explicit_empty_rejected(monkeypa is present as `{}` in the request body. The value matches the model default but `model_fields_set` distinguishes the two.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -17863,6 +18070,7 @@ async def test_update_key_non_admin_permissions_explicit_null_rejected(monkeypat """`_validate_update_key_data` rejects a non-admin when `permissions` is present as `null` in the request body.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -17892,6 +18100,7 @@ async def test_update_key_non_admin_omits_permissions_succeeds(monkeypatch): `permissions` is absent from the request body (personal-key fast path on an unrelated field).""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -17914,6 +18123,7 @@ async def test_update_key_admin_can_set_permissions(monkeypatch): """`_validate_update_key_data` accepts a PROXY_ADMIN caller for every shape of `permissions` in the request body.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -18120,6 +18330,7 @@ async def test_update_key_non_admin_disable_global_guardrails_rejected(monkeypat """`_validate_update_key_data` rejects a non-admin when `disable_global_guardrails` is true in the request body.""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -18148,6 +18359,7 @@ async def test_update_key_non_admin_resending_stored_disable_global_guardrails_a re-sends `metadata.disable_global_guardrails` that is already stored on the key (the Admin UI edit form round-trips the whole metadata JSON).""" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.jsonify_object = lambda data: data monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -18186,6 +18398,7 @@ async def test_regenerate_key_non_admin_disable_global_guardrails_rejected(monke existing_key = _make_regenerate_existing_key() mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_repo = MagicMock() mock_repo.table.find_unique = AsyncMock(return_value=existing_key) @@ -18398,6 +18611,7 @@ async def test_list_keys_rejects_invalid_expires(): from unittest.mock import Mock, patch mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): @@ -18423,6 +18637,7 @@ async def test_list_keys_forwards_expires_filter(expires_value, expected_forward from unittest.mock import Mock, patch mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) mock_user_info = LiteLLM_UserTable( user_id="admin-user", @@ -18462,6 +18677,7 @@ async def test_list_keys_without_expires_param_forwards_none(): from unittest.mock import Mock, patch mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) mock_user_info = LiteLLM_UserTable( user_id="admin-user", @@ -18525,6 +18741,7 @@ async def test_rotate_master_key_rotates_sso_identity_assertions( mock_prisma_client = AsyncMock() mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_tx = AsyncMock() mock_tx.litellm_proxymodeltable = MagicMock() @@ -18798,6 +19015,7 @@ def _estimate_key_row(token: str, metadata: dict): def _wire_update_key_fn(monkeypatch, existing_key): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db updated_key = MagicMock() updated_key.token = existing_key.token updated_key.key_alias = "my-alias" @@ -19220,8 +19438,10 @@ def _wire_key_generation_prisma(monkeypatch): created_key = MagicMock(token="hashed_token_123", litellm_budget_table=None, object_permission=None) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.insert_data = AsyncMock(return_value=created_key) mock_prisma_client.db = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken = MagicMock() mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -19466,6 +19686,7 @@ async def test_update_key_syncs_access_group_assigned_key_ids_in_both_directions } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=key_in_db ) @@ -19554,6 +19775,7 @@ async def test_update_key_leaves_access_groups_alone_when_field_is_unset(monkeyp } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=key_in_db ) @@ -19609,6 +19831,7 @@ async def test_bulk_update_keys_syncs_access_group_assigned_key_ids(monkeypatch) } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.update_data = AsyncMock(return_value={"data": {}}) _access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups) _setup_update_key_mocks(monkeypatch, mock_prisma_client) @@ -19672,6 +19895,7 @@ async def test_delete_key_withdraws_token_from_its_access_groups(monkeypatch): } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( return_value=[key_in_db] ) @@ -19723,6 +19947,7 @@ async def test_generate_key_records_token_in_its_access_groups(monkeypatch): created_key.updated_at = None mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.insert_data = AsyncMock(return_value=created_key) _access_group_table_mocks(monkeypatch, mock_prisma_client, access_groups) monkeypatch.setattr( @@ -19856,6 +20081,7 @@ async def test_key_write_paths_revoke_the_key_cache_before_syncing_access_groups } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=key_in_db ) @@ -19931,6 +20157,7 @@ async def test_update_key_syncs_many_access_groups_in_one_statement_per_directio } mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( return_value=key_in_db ) @@ -20299,6 +20526,7 @@ async def test_update_key_row_with_soft_budget_updates_budget_and_key_in_transac tx_context.__aenter__ = AsyncMock(return_value=tx) tx_context.__aexit__ = AsyncMock(return_value=None) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.tx.return_value = tx_context prisma_client.jsonify_object = lambda data: dict(data) @@ -20335,6 +20563,7 @@ async def test_update_key_row_with_soft_budget_propagates_transaction_error(): tx_context.__aenter__ = AsyncMock(return_value=tx) tx_context.__aexit__ = AsyncMock(return_value=None) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.tx.return_value = tx_context prisma_client.jsonify_object = lambda data: dict(data) @@ -20491,6 +20720,7 @@ async def test_key_creator_cannot_detach_project_without_admin_access(): token="project-detach-token", project_id="project-orbit", user_id="user-orbit", created_by="user-orbit", ) database: Final = MagicMock() + database.replica_db = database.db database.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing) with pytest.raises(HTTPException) as exc: await _validate_update_key_data( @@ -20687,6 +20917,7 @@ async def test_key_update_invalidates_cached_object_permission(monkeypatch): return row mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( side_effect=lambda **kwargs: _row(grants["old"]) ) @@ -20883,6 +21114,7 @@ async def test_key_update_evicts_object_permission_before_key_object(monkeypatch permission_id = "objperm-order" mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db existing_permission_row = MagicMock() existing_permission_row.model_dump.return_value = { "object_permission_id": permission_id, diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 5f7807650e1..4fe9b2d1ac2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -57,6 +57,7 @@ class MockPrismaClient: self.user_admin = user_admin self.sibling_deployments = sibling_deployments or [] self.db = self + self.replica_db = self.db async def find_unique(self, where): if self.team_exists: @@ -385,6 +386,7 @@ class TestModelManagementAuthChecks: ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: endpoint reads proxy server globals with no injection seam patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam @@ -515,6 +517,7 @@ class TestModelManagementAuthChecks: ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: endpoint reads proxy server globals with no injection seam patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: endpoint reads proxy server globals with no injection seam @@ -608,6 +611,7 @@ class TestModelManagementAuthChecks: existing_row = MagicMock() existing_row.model_dump.return_value = existing.model_dump() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() with ( @@ -702,6 +706,7 @@ class TestDeleteTeamModelAlias: # Create mock prisma client mock_prisma = MockPrismaClient(team_exists=True) mock_prisma.db = MockPrismaWrapper(model_aliases_list) + mock_prisma.replica_db = mock_prisma.db # Call the function await delete_team_model_alias( @@ -760,6 +765,7 @@ class TestDeleteTeamModelAlias: # Create mock prisma client mock_prisma = MockPrismaClient(team_exists=True) mock_prisma.db = MockPrismaWrapper(model_aliases_list) + mock_prisma.replica_db = mock_prisma.db # Call the function with non-existent model await delete_team_model_alias( @@ -790,6 +796,7 @@ class TestClearCache: ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_logging = MagicMock() with ( @@ -848,6 +855,7 @@ class TestClearCache: ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_logging = MagicMock() with ( @@ -1125,6 +1133,7 @@ class TestDeleteModelClearsRouterRegistry: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) @@ -1187,6 +1196,7 @@ class TestDeleteModelClearsRouterRegistry: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) @@ -1274,9 +1284,8 @@ class TestDeletedAutoRouterAvailability: find_unique=AsyncMock(return_value=row), delete=AsyncMock(return_value=row, side_effect=None if delete_succeeds else RuntimeError("delete failed")), ) - prisma = SimpleNamespace( - db=SimpleNamespace(litellm_proxymodeltable=table, query_raw=AsyncMock(return_value=[])) - ) + db = SimpleNamespace(litellm_proxymodeltable=table, query_raw=AsyncMock(return_value=[])) + prisma = SimpleNamespace(db=db, replica_db=db) monkeypatch.setattr(proxy_server, "prisma_client", prisma) monkeypatch.setattr(proxy_server, "store_model_in_db", True) admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) @@ -1374,6 +1383,7 @@ class TestUpdateModel: updated_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( return_value=existing_row ) @@ -1434,6 +1444,7 @@ class TestUpdateModel: updated_row = MagicMock() updated_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) mock_router = MagicMock() @@ -2529,6 +2540,7 @@ class TestAddAndDeleteModelLifecycle: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row) @@ -2642,6 +2654,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( @@ -2725,6 +2738,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( @@ -2802,6 +2816,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( @@ -2888,6 +2903,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( @@ -2970,6 +2986,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( @@ -3040,6 +3057,7 @@ class TestDeleteModelTeamAuth: ) mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( @@ -3159,6 +3177,7 @@ class TestDeleteModelTeamAuth: ) mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( @@ -3329,6 +3348,7 @@ class _TxPrismaClient: return False self.db = MagicMock() + self.replica_db = self.db self.db.tx = MagicMock(return_value=_TxCM()) @@ -4135,6 +4155,7 @@ class TestModelInfoServerDerivedPricingFilter: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row) @@ -4980,6 +5001,7 @@ class TestPatchModelBlockedAuthGate: existing_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( return_value=existing_row ) @@ -5023,6 +5045,7 @@ class TestPatchModelBlockedAuthGate: updated_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( return_value=existing_row ) @@ -5078,6 +5101,7 @@ class TestPatchModelRowDeletedBeforeWrite: existing_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( return_value=existing_row ) @@ -5462,6 +5486,7 @@ class TestDeleteEvictionsHoldTheReconcileLock: table.delete = AsyncMock(return_value=row) prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_proxymodeltable = table prisma.db.query_raw = AsyncMock(return_value=[]) @@ -5517,6 +5542,7 @@ class TestDeleteEvictionsHoldTheReconcileLock: tx_ctx.__aexit__ = AsyncMock(return_value=False) prisma = MagicMock() + prisma.replica_db = prisma.db prisma.db.tx = MagicMock(return_value=tx_ctx) monkeypatch.setattr( @@ -5898,6 +5924,7 @@ class TestStrategyRouterWriteValidation: admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -6043,6 +6070,7 @@ class TestStrategyRouterWriteValidation: self, db_models: list[str], existing_row: object = None, tuning_rows: list[dict[str, object]] | None = None ) -> None: self.db = self + self.replica_db = self.db self.tx_obj = TestStrategyRouterWriteValidation._FakeTx(db_models, tuning_rows=tuning_rows) self.litellm_proxymodeltable = MagicMock( create=AsyncMock(), update=AsyncMock(), find_unique=AsyncMock(return_value=existing_row) @@ -6681,6 +6709,7 @@ class TestStrategyRouterWriteValidation: } mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() @@ -7097,6 +7126,7 @@ class TestBlockModelResponseSerialization: updated_row = prisma_models.LiteLLM_ProxyModelTable(blocked=blocked, **row_fields) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) @@ -7180,6 +7210,7 @@ class TestAccessGroupModelSync: mock_prisma = MagicMock() mock_prisma.db = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(side_effect=query_raw) mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=row) @@ -7505,7 +7536,7 @@ class TestTeamMemberAutoRouterWrites: litellm_proxymodeltable=table, tx=MagicMock(return_value=context), ) - return MagicMock(db=db, transaction=transaction) + return MagicMock(db=db, replica_db=db, transaction=transaction) @staticmethod def _catalog() -> Router: @@ -7727,6 +7758,7 @@ class TestModelManagementActorEdges: actor: Final = UserAPIKeyAuth(user_id="internal-user", user_role=LitellmUserRoles.INTERNAL_USER) prisma: Final = MagicMock() + prisma.replica_db = prisma.db deployment: Final = Deployment( model_name="internal-model", litellm_params=LiteLLM_Params(model="openai/test-model"), @@ -7754,6 +7786,7 @@ class TestModelManagementActorEdges: user_id="view-only-user", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY ) prisma: Final = MagicMock() + prisma.replica_db = prisma.db deployment: Final = Deployment( model_name="view-only-model", litellm_params=LiteLLM_Params(model="openai/test-model"), @@ -7779,6 +7812,7 @@ class TestModelManagementActorEdges: actor: Final = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) prisma: Final = MagicMock() + prisma.replica_db = prisma.db deployment: Final = Deployment( model_name="database-disabled-model", litellm_params=LiteLLM_Params(model="openai/test-model"), @@ -7813,6 +7847,7 @@ class TestModelManagementActorEdges: updated_row: Final = MagicMock() updated_row.model_dump_json.return_value = "{}" prisma: Final = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) router: Final = MagicMock() @@ -7862,6 +7897,7 @@ class TestModelManagementActorEdges: updated_row: Final = MagicMock() updated_row.model_dump_json.return_value = "{}" prisma: Final = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) router: Final = MagicMock() @@ -7901,6 +7937,7 @@ class TestModelManagementActorEdges: model_id: Final = "config-model-id" prisma: Final = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) prisma.db.litellm_proxymodeltable.update = AsyncMock() router: Final = MagicMock() @@ -7944,6 +7981,7 @@ class TestModelManagementActorEdges: def test_post_model_new_binds_to_actor_guard(self): actor: Final = UserAPIKeyAuth(user_id="internal-user", user_role=LitellmUserRoles.INTERNAL_USER) prisma: Final = MagicMock() + prisma.replica_db = prisma.db with ( patch("litellm.proxy.proxy_server.prisma_client", prisma), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: [TQ008] route reads proxy-server state through its only test seam @@ -7983,6 +8021,7 @@ class TestModelManagementActorEdges: updated_by="admin", ) prisma: Final = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) router: Final = MagicMock() @@ -8024,6 +8063,7 @@ class TestModelManagementActorEdges: def test_patch_config_model_binds_to_patch_route(self): model_id: Final = "config-route-model-id" prisma: Final = MagicMock() + prisma.replica_db = prisma.db prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) prisma.db.litellm_proxymodeltable.update = AsyncMock() router: Final = MagicMock() diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index b066b3b80e6..626ee8f168c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -200,8 +200,10 @@ def _wire_team_delete_tx(prisma_client): # Mock prisma_client mock_prisma_client = MagicMock() +mock_prisma_client.replica_db = mock_prisma_client.db # Set up async mock for db operations mock_prisma_client.db = MagicMock() +mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable = MagicMock() mock_prisma_client.db.litellm_teamtable.update = AsyncMock() mock_prisma_client.db.litellm_auditlog = MagicMock() @@ -583,6 +585,7 @@ async def test_new_team_rejects_a_duration_that_never_advances( from litellm.proxy.management_endpoints.team_endpoints import new_team mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db mock_team_create = AsyncMock() mock_db_client.db.litellm_teamtable = MagicMock() mock_db_client.db.litellm_teamtable.create = mock_team_create @@ -612,6 +615,7 @@ async def test_update_team_rejects_a_duration_that_never_advances( from litellm.proxy.management_endpoints.team_endpoints import update_team mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db mock_find_unique = AsyncMock(return_value=None) mock_db_client.db.litellm_teamtable = MagicMock() mock_db_client.db.litellm_teamtable.find_unique = mock_find_unique @@ -641,6 +645,7 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth): mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db # Mock object permission table creation mock_object_perm_create = AsyncMock( @@ -719,6 +724,7 @@ async def test_new_team_persists_tpd_limit(mock_db_client, mock_admin_auth): mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db mock_db_client.db.litellm_modeltable = MagicMock() mock_db_client.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) @@ -764,6 +770,7 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db # Track what data is passed to object permission create created_permission_data = {} @@ -882,6 +889,7 @@ async def test_new_team_disable_auto_add_proxy_admin_flag( mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db team_create_result = MagicMock(team_id="team-789") team_create_result.model_dump.return_value = {"team_id": "team-789"} @@ -942,6 +950,7 @@ async def test_team_update_object_permissions_existing_permission(monkeypatch): # Mock prisma client mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Mock existing team with object_permission_id @@ -1014,6 +1023,7 @@ async def test_team_update_object_permissions_no_existing_permission(monkeypatch # Mock prisma client mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) existing_team_row_no_perm = LiteLLM_TeamTable( @@ -1074,6 +1084,7 @@ async def test_team_update_object_permissions_missing_permission_record(monkeypa # Mock prisma client mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) existing_team_row_missing_perm = LiteLLM_TeamTable( @@ -1197,6 +1208,7 @@ async def test_add_team_member_budget_table_success(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock budget record mock_budget_record = MagicMock() @@ -1244,6 +1256,7 @@ async def test_add_team_member_budget_table_exception_handling(): # Mock prisma client to raise an exception mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock( side_effect=Exception("Database connection failed") ) @@ -1297,6 +1310,7 @@ async def test_add_team_member_budget_table_budget_not_found(): # Mock prisma client to return None (budget not found) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) # Create team info response object @@ -1803,6 +1817,7 @@ async def test_process_team_members_single_member(): # Mock dependencies mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_team = MagicMock(spec=LiteLLM_TeamTable) mock_team.metadata = {"team_member_budget_id": "budget-123"} mock_team.default_team_member_models = None @@ -1864,6 +1879,7 @@ async def test_process_team_members_multiple_members(): # Mock dependencies mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_team = MagicMock(spec=LiteLLM_TeamTable) mock_team.metadata = None mock_team.default_team_member_models = None @@ -2034,6 +2050,7 @@ async def test_add_team_members_reconciles_against_freshly_locked_row(): tx_cm.__aexit__ = AsyncMock(return_value=None) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.tx = MagicMock(return_value=tx_cm) with patch( @@ -2107,6 +2124,7 @@ async def test_add_team_members_runs_member_writes_on_the_lock_holding_transacti tx_cm.__aexit__ = AsyncMock(return_value=None) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.tx = MagicMock(return_value=tx_cm) type(prisma_client).db = PropertyMock( side_effect=AssertionError("member writes must not reach for a second pooled connection") @@ -2172,6 +2190,7 @@ async def test_add_team_members_skips_budget_and_membership_writes_for_members_a tx_cm.__aexit__ = AsyncMock(return_value=None) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.tx = MagicMock(return_value=tx_cm) _, updated_users, updated_team_memberships = await _add_team_members_to_team( @@ -2220,6 +2239,7 @@ async def test_add_team_members_writes_nothing_when_the_team_is_deleted_mid_requ tx_cm.__aexit__ = AsyncMock(return_value=None) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.tx = MagicMock(return_value=tx_cm) prisma_client.db.execute_raw = AsyncMock() prisma_client.db.litellm_teammembership.delete_many = AsyncMock() @@ -2340,6 +2360,7 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name): new_callable=AsyncMock, ) as mock_cache_team, ): + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( return_value=existing_team ) @@ -2417,6 +2438,7 @@ async def test_team_model_add_delete_keep_model_aliases_in_team_cache(endpoint_n return SimpleNamespace(team_id="team-1234", model_dump=lambda: row) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(model_dump=lambda: columns)) prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=update) prisma_client.db.execute_raw = AsyncMock(return_value=None) @@ -2524,6 +2546,7 @@ async def test_team_write_404s_when_row_vanishes_before_update(endpoint_name): return_value=existing_team, ), ): + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( return_value=existing_team ) @@ -2574,6 +2597,7 @@ async def test_update_team_team_member_budget_not_passed_to_db( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" ) as mock_upsert_budget, ): + mock_prisma_client.replica_db = mock_prisma_client.db # Setup mock prisma client mock_existing_team = MagicMock() mock_existing_team.model_dump.return_value = { @@ -3142,6 +3166,7 @@ async def test_update_team_with_team_member_budget_duration( "litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table" ) as mock_upsert_budget, ): + mock_prisma_client.replica_db = mock_prisma_client.db mock_existing_team = MagicMock() mock_existing_team.model_dump.return_value = { "team_id": "test_team_id", @@ -3230,6 +3255,7 @@ async def test_backfill_team_member_budget_entries_creates_missing_memberships() existing_membership.user_id = "user-A" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teammembership.find_many = AsyncMock( return_value=[existing_membership] ) @@ -3305,6 +3331,7 @@ async def test_backfill_team_member_budget_entries_no_op_when_all_exist(): existing_b.user_id = "user-B" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teammembership.find_many = AsyncMock( return_value=[existing_a, existing_b] ) @@ -3352,6 +3379,7 @@ async def test_backfill_team_member_budget_entries_populates_null_budget_id_on_e existing_b.user_id = "user-B" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teammembership.find_many = AsyncMock( return_value=[existing_a, existing_b] ) @@ -3388,6 +3416,7 @@ async def test_backfill_team_member_budget_entries_empty_members(): ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teammembership.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teammembership.create_many = AsyncMock(return_value=None) @@ -3599,6 +3628,7 @@ async def test_bulk_team_member_add_all_users_flag(): return_value=mock_team_response, ) as mock_team_member_add, ): + mock_prisma.replica_db = mock_prisma.db # Mock the database find_many call mock_prisma.db.litellm_usertable.find_many = AsyncMock( return_value=mock_db_users @@ -3725,6 +3755,7 @@ async def test_list_team_v2_security_check_non_admin_user(): return_value=None, ), ): + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.return_value = MagicMock() # Mock non-None prisma client # Should raise HTTPException with 401 status @@ -3775,6 +3806,7 @@ async def test_list_team_v2_security_check_non_admin_user_other_user(): return_value=None, ), ): + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.return_value = MagicMock() # Mock non-None prisma client # Should raise HTTPException with 401 status @@ -3818,9 +3850,11 @@ async def test_list_team_v2_security_check_non_admin_user_own_teams(): patch("litellm.proxy.proxy_server.user_api_key_cache"), patch("litellm.proxy.proxy_server.proxy_logging_obj"), ): + mock_prisma_client.replica_db = mock_prisma_client.db # Mock prisma client and database operations mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db # Mock get_user_object to return a user with teams from litellm.proxy._types import LiteLLM_UserTable @@ -3883,9 +3917,11 @@ async def test_list_team_v2_security_check_admin_user(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: + mock_prisma_client.replica_db = mock_prisma_client.db # Mock prisma client and database operations mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db # Mock team lookup mock_teams = [ @@ -3934,9 +3970,11 @@ async def test_list_team_v2_with_status_deleted(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: + mock_prisma_client.replica_db = mock_prisma_client.db # Mock prisma client and database operations mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db # Mock deleted teams mock_deleted_team1 = Mock( @@ -4030,8 +4068,10 @@ async def test_list_team_v2_includes_litellm_model_table(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: # test-quality-ok: this file's DB-mock convention + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db mock_db.litellm_teamtable.find_many = AsyncMock( side_effect=lambda **kw: [_team_row("team_1", kw.get("include"))] @@ -4118,8 +4158,10 @@ async def test_list_team_v2_org_admin_sees_org_teams(): return_value=mock_user, ), ): + mock_prisma.replica_db = mock_prisma.db mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_team = Mock() mock_team.model_dump.return_value = { @@ -4209,8 +4251,10 @@ async def test_list_team_v2_org_admin_own_user_id_sees_all_org_teams(): return_value=mock_user, ), ): + mock_prisma.replica_db = mock_prisma.db mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_team_1 = Mock() mock_team_1.model_dump.return_value = { @@ -4352,6 +4396,7 @@ async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs( return len(await find_many(where)) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many) prisma_client.db.litellm_teamtable.count = AsyncMock(side_effect=count) prisma_client.db.litellm_verificationtoken.group_by = AsyncMock(return_value=[]) @@ -4447,6 +4492,7 @@ async def test_list_team_v1_org_admin_own_query_keeps_memberships_in_other_orgs( return [t for t in all_teams if t.organization_id in where["organization_id"]["in"]] prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many) async def list_teams(user_id): @@ -4514,7 +4560,9 @@ async def test_list_team_v2_org_admin_cannot_view_other_orgs(): return_value=mock_user, ), ): + mock_prisma.replica_db = mock_prisma.db mock_prisma.db = Mock() + mock_prisma.replica_db = mock_prisma.db with pytest.raises(HTTPException) as exc_info: await list_team_v2( @@ -4604,8 +4652,10 @@ async def test_list_team_v2_org_admin_with_user_id_returns_user_teams(): side_effect=mock_get_user_object, ), ): + mock_prisma.replica_db = mock_prisma.db mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_team = Mock() mock_team.model_dump.return_value = { @@ -4661,6 +4711,7 @@ async def test_list_team_v2_with_invalid_status(): ) mock_prisma_client = Mock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock prisma_client to be non-None with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): @@ -4700,8 +4751,10 @@ async def test_list_team_v2_search_builds_or_clause(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[]) mock_db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4747,8 +4800,10 @@ async def test_list_team_v2_search_team_id_match_prefix(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[]) mock_db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4816,8 +4871,10 @@ async def test_list_team_v2_search_composes_with_user_id_filter(): new=AsyncMock(return_value=None), ), ): + mock_prisma.replica_db = mock_prisma.db mock_db = Mock() mock_prisma.db = mock_db + mock_prisma.replica_db = mock_prisma.db mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[]) mock_db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4864,8 +4921,10 @@ async def test_list_team_v2_populates_keys_count(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db team_a = Mock() team_a.team_id = "team_a" @@ -4931,8 +4990,10 @@ async def test_list_team_v2_keys_count_skipped_for_empty_page(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[]) mock_db.litellm_teamtable.count = AsyncMock(return_value=0) @@ -4972,8 +5033,10 @@ async def test_list_team_v2_keys_count_skipped_for_deleted_status(): ) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client: + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = Mock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db mock_deleted = Mock() mock_deleted.team_id = "team_d" @@ -5234,6 +5297,7 @@ async def test_delete_team_writes_deleted_audit_log_for_team_keys( team_key = LiteLLM_VerificationToken(token="hashed-doomed-key", team_id="team-doomed") mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) mock_prisma_client.delete_data = AsyncMock(return_value={"deleted_teams": ["team-doomed"]}) mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock() @@ -5883,6 +5947,7 @@ async def test_new_team_max_budget_exceeds_user_max_budget(): "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -5952,6 +6017,7 @@ async def test_new_team_max_budget_within_user_limit(): "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -6087,6 +6153,7 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): "litellm.proxy.management_endpoints.team_endpoints.get_org_object" ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -6233,6 +6300,7 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): "litellm.proxy.management_endpoints.team_endpoints.get_org_object" ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -6374,6 +6442,7 @@ async def test_new_team_standalone_validates_against_user_models(monkeypatch): "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -6443,6 +6512,7 @@ async def test_new_team_standalone_validates_against_user_budget(): "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Setup basic mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -6519,6 +6589,7 @@ async def test_new_team_org_scoped_budget_exceeds_org_limit(): "litellm.proxy.management_endpoints.team_endpoints.get_org_object" ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -6598,6 +6669,7 @@ async def test_new_team_org_scoped_models_not_in_org_models(): "litellm.proxy.management_endpoints.team_endpoints.get_org_object" ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -6672,6 +6744,7 @@ async def test_update_team_standalone_budget_raise_blocked_for_team_admin(): "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ), ): + mock_prisma.replica_db = mock_prisma.db mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-123" mock_existing_team.organization_id = None # Standalone team @@ -6740,6 +6813,7 @@ async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ), ): + mock_prisma.replica_db = mock_prisma.db mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-123" mock_existing_team.organization_id = None @@ -6829,6 +6903,7 @@ async def test_update_team_standalone_budget_removal_blocked_for_team_admin(): "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ), ): + mock_prisma.replica_db = mock_prisma.db mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-123" mock_existing_team.organization_id = None @@ -6899,6 +6974,7 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ), ): + mock_prisma.replica_db = mock_prisma.db _TeamRowStore( mock_prisma.db.litellm_teamtable, { @@ -6972,6 +7048,7 @@ async def test_update_team_standalone_unchanged_budget_allowed( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing standalone team (no organization_id) with budget=$500 mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-unchanged-budget-123" @@ -7071,6 +7148,7 @@ async def test_update_team_standalone_lower_budget_allowed( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db _TeamRowStore( mock_prisma.db.litellm_teamtable, { @@ -7157,6 +7235,7 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit(): new=AsyncMock(return_value=mock_org), ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-456" @@ -7234,6 +7313,7 @@ async def test_update_team_standalone_models_not_gated_by_user_limit( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing standalone team (no organization_id) mock_existing_team = MagicMock() mock_existing_team.team_id = "standalone-team-models-123" @@ -7341,6 +7421,7 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit( new=AsyncMock(return_value=mock_org), ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-budget-123" @@ -7452,6 +7533,7 @@ async def test_update_team_org_scoped_models_bypasses_user_limit( new=AsyncMock(return_value=mock_org), ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-models-123" @@ -7556,6 +7638,7 @@ async def test_update_team_org_scoped_models_not_in_org_models(): new=AsyncMock(return_value=mock_org), ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-models-fail-123" @@ -7647,6 +7730,7 @@ async def test_update_team_org_scoped_models_with_all_proxy_models( new=AsyncMock(return_value=mock_org), ) as mock_get_org, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-all-proxy-models-123" @@ -7753,6 +7837,7 @@ async def test_update_team_tpm_limit_not_gated_by_user_limit( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ), ): + mock_prisma.replica_db = mock_prisma.db # Mock existing standalone team mock_existing_team = MagicMock() mock_existing_team.team_id = "team-tpm-test-123" @@ -7836,6 +7921,7 @@ async def test_update_team_rpm_limit_not_gated_by_user_limit( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ), ): + mock_prisma.replica_db = mock_prisma.db # Mock existing standalone team mock_existing_team = MagicMock() mock_existing_team.team_id = "team-rpm-test-123" @@ -7898,6 +7984,7 @@ async def test_update_team_persists_tpd_limit(disable_audit_logging_for_mocked_t "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ), ): + mock_prisma.replica_db = mock_prisma.db existing_team = MagicMock(team_id="team-tpd", organization_id=None, model_id=None, tpd_limit=None) existing_team.model_dump.return_value = {"team_id": "team-tpd", "organization_id": None} mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) @@ -7977,6 +8064,7 @@ async def test_new_team_org_scoped_tpm_exceeds_org_limit(): new=AsyncMock(return_value=mock_org), ), ): + mock_prisma.replica_db = mock_prisma.db mock_license.is_team_count_over_limit.return_value = False mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_prisma.get_data = AsyncMock(return_value=None) @@ -8052,6 +8140,7 @@ async def test_new_team_org_scoped_rpm_exceeds_org_limit(): new=AsyncMock(return_value=mock_org), ), ): + mock_prisma.replica_db = mock_prisma.db mock_license.is_team_count_over_limit.return_value = False mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_prisma.get_data = AsyncMock(return_value=None) @@ -8137,6 +8226,7 @@ async def test_new_team_org_scoped_tpm_rpm_bypasses_user_limit(): new=AsyncMock(), ), ): + mock_prisma.replica_db = mock_prisma.db mock_license.is_team_count_over_limit.return_value = False mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_prisma.get_data = AsyncMock(return_value=None) @@ -8236,6 +8326,7 @@ async def test_update_team_org_scoped_tpm_exceeds_org_limit(): new=AsyncMock(return_value=mock_org), ), ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-tpm-123" @@ -8324,6 +8415,7 @@ async def test_update_team_org_scoped_rpm_exceeds_org_limit(): new=AsyncMock(return_value=mock_org), ), ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-rpm-123" @@ -8418,6 +8510,7 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit( new=AsyncMock(return_value=mock_org), ), ): + mock_prisma.replica_db = mock_prisma.db # Mock existing org-scoped team mock_existing_team = MagicMock() mock_existing_team.team_id = "org-team-update-bypass-123" @@ -8550,6 +8643,7 @@ async def test_update_team_guardrails_with_org_id( ), patch("litellm.proxy.proxy_server.llm_router", MagicMock()), ): + mock_prisma.replica_db = mock_prisma.db # Mock existing team - must have compatible models with organization mock_existing_team = MagicMock() mock_existing_team.team_id = "team-guardrails-123" @@ -8732,6 +8826,7 @@ def test_transform_teams_to_deleted_records_empty_list(): @pytest.mark.asyncio async def test_save_deleted_team_records(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_create_many = AsyncMock() mock_prisma_client.db.litellm_deletedteamtable.create_many = mock_create_many @@ -8758,6 +8853,7 @@ async def test_save_deleted_team_records(): @pytest.mark.asyncio async def test_save_deleted_team_records_empty_list(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_create_many = AsyncMock() mock_prisma_client.db.litellm_deletedteamtable.create_many = mock_create_many @@ -8769,6 +8865,7 @@ async def test_save_deleted_team_records_empty_list(): @pytest.mark.asyncio async def test_persist_deleted_team_records(): mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_create_many = AsyncMock() mock_prisma_client.db.litellm_deletedteamtable.create_many = mock_create_many @@ -8814,6 +8911,7 @@ async def test_delete_team_persists_deleted_teams( from litellm.proxy._types import DeleteTeamRequest mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_dict = UserAPIKeyAuth( user_id="admin-user", api_key="sk-admin", @@ -8932,6 +9030,7 @@ async def test_delete_team_sweeps_references_outside_members_with_roles( return 1 mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=doomed_team) mock_prisma_client.delete_data = AsyncMock(return_value={"deleted_keys": 0}) mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock() @@ -9040,6 +9139,7 @@ async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes( team_key = LiteLLM_VerificationToken(token="hashed-doomed-key", team_id="team-doomed") mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) mock_prisma_client.delete_data = AsyncMock(return_value={"deleted_teams": ["team-doomed"]}) mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock() @@ -9107,6 +9207,7 @@ async def test_delete_team_failing_locked_sweep_rolls_back_the_delete_and_leaves ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) mock_prisma_client.delete_data = AsyncMock(return_value={"deleted_teams": ["team-doomed"]}) mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock() @@ -9174,6 +9275,7 @@ async def test_delete_team_broadcasts_cache_invalidation_to_other_workers( ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) mock_prisma_client.delete_data = AsyncMock(return_value={"deleted_teams": ["team-doomed"]}) mock_prisma_client.db.litellm_deletedteamtable.create_many = AsyncMock() @@ -9242,6 +9344,7 @@ async def test_delete_team_survives_a_failing_cache_backend( ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) mock_delete_data = AsyncMock(return_value={"deleted_teams": ["team-doomed"]}) mock_prisma_client.delete_data = mock_delete_data @@ -9303,6 +9406,7 @@ async def test_team_member_delete_persists_deleted_keys(monkeypatch): ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_user_api_key_dict = UserAPIKeyAuth( user_id="admin-user", api_key="sk-admin", @@ -9466,6 +9570,7 @@ async def test_team_member_delete_evicts_jwt_key_mapping_cache_of_the_keys_it_de key1 = LiteLLM_VerificationToken(token="hashed-token-1", user_id="user-123", team_id="team-1") mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) mock_prisma_client.db.litellm_usertable.find_many = AsyncMock( return_value=[MagicMock(user_id="user-123", teams=["team-1"])] @@ -9527,6 +9632,7 @@ async def test_delete_team_evicts_jwt_key_mapping_cache_of_the_keys_it_deletes( ) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) async def cascading_delete_data(team_id_list, table_name): @@ -9703,6 +9809,7 @@ async def test_new_team_soft_budget_validation( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Setup mocks mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False @@ -9905,6 +10012,7 @@ async def test_update_team_soft_budget_validation( "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() ) as mock_audit, ): + mock_prisma.replica_db = mock_prisma.db # Mock existing team with existing budgets mock_existing_team = MagicMock() mock_existing_team.team_id = "test-team-123" @@ -10021,6 +10129,7 @@ async def test_new_team_with_router_settings(mock_db_client, mock_admin_auth): mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db mock_db_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[ SimpleNamespace(model_id="weighted-id", model_name="group", model_info={}) ]) @@ -10294,6 +10403,7 @@ async def test_update_team_with_router_settings( # Configure mocked prisma client mock_db_client.jsonify_team_object = lambda db_data: db_data mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db mock_db_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[ SimpleNamespace(model_id="weighted-id", model_name="group", model_info={}) ]) @@ -10560,6 +10670,7 @@ async def test_validate_and_populate_member_user_info_both_provided_match(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock user object that matches both email and user_id mock_user = MagicMock() @@ -10598,6 +10709,7 @@ async def test_validate_and_populate_member_user_info_only_email_provided(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock user object from find_first mock_user_find_first = MagicMock() @@ -10647,6 +10759,7 @@ async def test_validate_and_populate_member_user_info_only_user_id_not_found(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock find_unique to return None (user not found) mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) @@ -10751,6 +10864,7 @@ async def test_list_team_v1_batches_key_queries(): return_value=[], ), ): + mock_prisma_client.replica_db = mock_prisma_client.db async def filtered_find_many(**kwargs): where = kwargs.get("where", {}) @@ -10860,6 +10974,7 @@ class TestBatchResolveAccessGroupResources: fake_row.access_agent_ids = ["agent-1", "agent-2"] fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( return_value=[fake_row] ) @@ -10891,6 +11006,7 @@ class TestBatchResolveAccessGroupResources: row2.access_agent_ids = ["agent-2"] fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( return_value=[row1, row2] ) @@ -10915,6 +11031,7 @@ class TestBatchResolveAccessGroupResources: row1.access_agent_ids = [] fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( return_value=[row1] ) @@ -10952,6 +11069,7 @@ class TestBatchResolveAccessGroupResources: fake_find_many = AsyncMock(return_value=[row1]) fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_accessgrouptable.find_many = fake_find_many with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma): @@ -10994,6 +11112,7 @@ class TestResolveTeamAccessGroupResources: row2.access_agent_ids = ["agent-1"] fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_accessgrouptable.find_many = AsyncMock( return_value=[row1, row2] ) @@ -11099,6 +11218,7 @@ async def test_update_team_rejects_unauthorized_caller(): return_value=False, ), ): + mock_prisma_client.replica_db = mock_prisma_client.db mock_existing_team = MagicMock() mock_existing_team.model_dump.return_value = { "team_id": "team-123", @@ -11447,6 +11567,7 @@ async def test_new_team_encrypts_callback_vars( ) mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db mock_db_client.db.litellm_teamtable = MagicMock() team_create_result = MagicMock(team_id="team-456", object_permission_id=None) team_create_result.model_dump.return_value = {"team_id": "team-456"} @@ -12005,6 +12126,7 @@ async def test_team_info_forwards_key_limit_to_get_data(): from litellm.proxy.management_endpoints import team_endpoints mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( return_value=LiteLLM_TeamTable(team_id="team-1") ) @@ -12047,6 +12169,7 @@ async def test_team_info_returns_model_aliases(): ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.get_data = AsyncMock(return_value=[]) @@ -12093,6 +12216,7 @@ async def test_team_info_hydrates_member_names_and_emails_from_the_user_table(): ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.get_data = AsyncMock(return_value=[]) @@ -12132,6 +12256,7 @@ async def test_update_model_table_clears_aliases_with_empty_map(): leaves the model table untouched. """ mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_modeltable.create = AsyncMock() mock_prisma.db.litellm_modeltable.upsert = AsyncMock( return_value=MagicMock(id="model-123") @@ -12254,6 +12379,7 @@ async def test_new_team_rejects_reserved_ui_session_team_id(): patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server._license_check") as mock_license, ): + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_license.is_team_count_over_limit.return_value = False mock_prisma.get_data = AsyncMock(return_value=None) @@ -12350,6 +12476,7 @@ async def _drive_team_write( new=AsyncMock(), ), ): + pc.replica_db = pc.db pc.db.litellm_teamtable.find_unique = AsyncMock( return_value=None if find_returns_none else existing ) @@ -12821,6 +12948,7 @@ async def test_new_team_validator_runs_without_metadata_and_rejection_blocks_cre patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), _configured_team_metadata_validator(validator), ): + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_prisma.db.litellm_teamtable.create = AsyncMock() _wire_team_create_tx(mock_prisma) @@ -12853,6 +12981,7 @@ async def test_new_team_validator_accept_proceeds_to_create(mock_db_client, mock mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db team_create_result = MagicMock(team_id="team-accept-1") team_create_result.model_dump.return_value = {"team_id": "team-accept-1"} @@ -12896,6 +13025,7 @@ async def test_new_team_rejection_precedes_model_alias_write(): patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), _configured_team_metadata_validator(validator), ): + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_prisma.db.litellm_teamtable.create = AsyncMock() _wire_team_create_tx(mock_prisma) @@ -13052,6 +13182,7 @@ async def test_get_all_team_memberships_validates_rows(): } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership_row]) result = await get_all_team_memberships(mock_prisma_client, ["team-1"], user_id="member-1") @@ -13085,6 +13216,7 @@ async def test_list_available_teams_filters_joined_and_validates_rows(monkeypatc open_team_row.model_dump = lambda: {"team_id": "team-open", "team_alias": "open team"} mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[open_team_row]) @@ -13278,6 +13410,7 @@ async def test_resolve_existing_member_user_ids_matches_caller_supplied_user_ids ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db find_many = AsyncMock( return_value=[LiteLLM_UserTable(user_id="by-id", max_budget=None, spend=0.0, user_email=None, models=[])] ) @@ -13517,6 +13650,7 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp member = Member(user_email="invitee@example.com", role="user") mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id") @@ -13620,6 +13754,7 @@ async def test_new_team_created_audit_event_carries_the_final_roster(monkeypatch audit_logger = _wire_audit_log_callback(monkeypatch) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) mock_prisma.jsonify_team_object = lambda db_data: db_data mock_prisma.get_data = AsyncMock(return_value=None) @@ -13740,6 +13875,7 @@ async def test_team_member_update_role_change_emits_a_roster_audit_event(monkeyp audit_logger = _wire_audit_log_callback(monkeypatch) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db team_row = LiteLLM_TeamTable( team_id="team-role-audit", team_alias="role-audit", @@ -13856,6 +13992,7 @@ async def test_team_member_update_role_change_rewrites_the_roster_it_read_under_ ) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=_roster_writer(team_row)) mock_prisma_client.db.litellm_auditlog.create = AsyncMock() @@ -13899,6 +14036,7 @@ async def test_team_member_update_role_change_404s_when_the_team_is_gone_under_t members_with_roles=[Member(user_id="bob", role="user")], ) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=[snapshot, None]) mock_prisma_client.db.litellm_teamtable.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -13935,6 +14073,7 @@ async def test_team_member_update_role_change_404s_when_the_member_left_before_t members_with_roles=[Member(user_id="alice", role="admin")], ) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=[snapshot, locked_row]) mock_prisma_client.db.litellm_teamtable.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -14076,6 +14215,7 @@ async def test_member_add_audit_reports_only_the_users_it_created_plus_the_roste audit_logger = _wire_audit_log_callback(monkeypatch) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_auditlog.create = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -14128,6 +14268,7 @@ async def test_delete_team_emits_only_the_deleted_audit_event(monkeypatch): model_spend={}, ) mock_prisma = AsyncMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) mock_prisma.get_data = AsyncMock( return_value=SimpleNamespace(json=lambda **_kwargs: team.model_dump_json(exclude_none=True)) @@ -14354,6 +14495,7 @@ def _wire_update_team(stack, existing_metadata): from unittest.mock import AsyncMock, MagicMock, patch mock_prisma_client = stack.enter_context(patch("litellm.proxy.proxy_server.prisma_client")) + mock_prisma_client.replica_db = mock_prisma_client.db stack.enter_context(patch("litellm.proxy.proxy_server.llm_router")) stack.enter_context(patch("litellm.proxy.proxy_server.user_api_key_cache")) stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) @@ -14650,6 +14792,7 @@ def _wire_new_team_prisma(mock_db_client): mock_db_client.jsonify_team_object = lambda db_data: db_data mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db created_team = MagicMock(team_id="team-defaults") created_team.model_dump.return_value = {"team_id": "team-defaults"} @@ -14889,6 +15032,7 @@ async def test_update_team_syncs_access_group_assigned_team_ids_in_both_directio new_callable=AsyncMock, ) as invalidate_cache, ): + prisma.replica_db = prisma.db prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) prisma.db.litellm_teamtable.update = AsyncMock(return_value=updated_team) prisma.db.tx = fake_db.tx @@ -14963,7 +15107,7 @@ async def test_sync_reads_the_committed_team_row_rather_than_the_callers_snapsho access_groups = {"ag-1": ["team-a", "team-b"], "ag-2": ["team-a"], "ag-3": []} teams = {"team-a": ["ag-2", "ag-3"]} fake_db = _FakeMirrorDb(access_groups, teams, plain_lists=True) - prisma_client = SimpleNamespace(db=SimpleNamespace(tx=fake_db.tx)) + prisma_client = SimpleNamespace(db=SimpleNamespace(tx=fake_db.tx), replica_db=SimpleNamespace(tx=fake_db.tx)) with patch( "litellm.proxy.management_helpers.access_group_team_sync.invalidate_access_group_cache", @@ -15021,6 +15165,7 @@ async def test_new_team_and_delete_team_both_drive_the_mirror( new_callable=AsyncMock, ) as invalidate_cache, ): + prisma.replica_db = prisma.db prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) prisma.db.litellm_teamtable.count = AsyncMock(return_value=0) prisma.db.tx = fake_db.tx @@ -15050,6 +15195,7 @@ async def test_new_team_and_delete_team_both_drive_the_mirror( new_callable=AsyncMock, ) as sync, ): + prisma.replica_db = prisma.db prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) prisma.delete_data = AsyncMock(return_value=[team_row]) @@ -15174,6 +15320,7 @@ async def test_reset_team_member_spend_fn_success(monkeypatch): from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_proxy_logging_obj = MagicMock() real_cache = UserApiKeyCache() await real_cache.async_set_cache(key="team-1_member-1", value="stale-membership") @@ -15223,6 +15370,7 @@ async def test_reset_team_member_spend_fn_success(monkeypatch): @pytest.mark.asyncio async def test_reset_team_member_spend_fn_membership_not_found(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) @@ -15247,6 +15395,7 @@ async def test_reset_team_member_spend_fn_membership_not_found(monkeypatch): @pytest.mark.asyncio async def test_reset_team_member_spend_fn_team_not_found(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) @@ -15272,6 +15421,7 @@ async def test_reset_team_member_spend_fn_forbidden_for_non_admin(monkeypatch): """A caller who is neither proxy admin, org admin, nor this team's admin must be refused, matching every other team-mutating endpoint's authorization.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) @@ -15299,6 +15449,7 @@ async def test_reset_team_member_spend_fn_team_admin_cannot_reset_own_spend(monk and repeatedly zero it right before it crosses their per-member cap, consuming the shared team budget without the configured limit ever binding (Veria finding on PR #37971).""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) @@ -15329,6 +15480,7 @@ async def test_reset_team_member_spend_fn_proxy_admin_can_reset_own_spend(monkey """The self-reset guard is scoped to non-proxy-admin roles: a proxy admin resetting their own membership spend is the platform-wide trust boundary, not a team-scoped one.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) @@ -15365,6 +15517,7 @@ async def test_reset_team_member_budget_fn_relinks_custom_member_to_team_default from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db real_cache = UserApiKeyCache() await real_cache.async_set_cache(key="team-1_member-1", value="stale-membership") await real_cache.async_set_cache(key="team_membership:member-1:team-1", value="stale-membership") @@ -15414,6 +15567,7 @@ async def test_reset_team_member_budget_fn_detaches_member_when_team_has_no_usab monkeypatch, team_obj, default_row ): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db membership_row = LiteLLM_TeamMembership(user_id="member-1", team_id="team-1", budget_id="custom-b1") mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) mock_prisma_client.db.litellm_teammembership.update = AsyncMock(return_value=membership_row) @@ -15442,6 +15596,7 @@ async def test_reset_team_member_budget_fn_detaches_member_when_team_has_no_usab @pytest.mark.asyncio async def test_reset_team_member_budget_fn_membership_not_found(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_teammembership.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -15463,6 +15618,7 @@ async def test_reset_team_member_budget_fn_membership_not_found(monkeypatch): @pytest.mark.asyncio async def test_reset_team_member_budget_fn_forbidden_for_non_admin(monkeypatch): mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_teammembership.update = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()) @@ -15494,6 +15650,7 @@ async def _team_info_budget_sources( from litellm.proxy.management_endpoints import team_endpoints mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.db.litellm_budgettable.find_unique = AsyncMock(return_value=default_budget_row) mock_prisma.get_data = AsyncMock(return_value=[]) @@ -15581,6 +15738,7 @@ async def test_team_member_update_invalidates_team_member_spend_state_when_budge from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db real_cache = UserApiKeyCache() await real_cache.async_set_cache(key="team-1_member-1", value="stale-membership") await real_cache.async_set_cache(key="team_membership:member-1:team-1", value="stale-membership") @@ -15635,6 +15793,7 @@ async def test_team_member_update_skips_invalidation_when_no_budget_fields_sent( from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db real_cache = UserApiKeyCache() await real_cache.async_set_cache(key="team-1_member-1", value="still-fresh-membership") real_spend_counter_cache = DualCache() @@ -15947,6 +16106,7 @@ async def test_team_info_returns_parent_organization_models(organization, expect ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.get_data = AsyncMock(return_value=[]) @@ -15997,6 +16157,7 @@ async def test_team_info_reports_parent_organization_models_only_to_team_manager ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) mock_prisma.get_data = AsyncMock(return_value=[]) @@ -16114,6 +16275,7 @@ async def test_new_team_persists_model_max_budget(mock_db_client, mock_admin_aut mock_db_client.get_data = AsyncMock(return_value=None) mock_db_client.update_data = AsyncMock(return_value=MagicMock()) mock_db_client.db = MagicMock() + mock_db_client.replica_db = mock_db_client.db mock_db_client.db.litellm_modeltable = MagicMock() mock_db_client.db.litellm_modeltable.create = AsyncMock(return_value=MagicMock(id="model123")) @@ -16201,6 +16363,7 @@ async def test_update_team_clearing_model_max_budget_writes_an_empty_mapping( patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), # test-quality-ok: proxy_server module global is the endpoint's only injection point patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point ): + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( return_value=_existing_team_with_model_caps(_EXISTING_TEAM_MODEL_CAPS) ) @@ -16234,6 +16397,7 @@ async def test_update_team_model_max_budget_raise_blocked_for_team_admin(): patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point patch("litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock()), # test-quality-ok: stubs the audit write so the test observes only the team update result ): + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( return_value=_existing_team_with_model_caps(_EXISTING_TEAM_MODEL_CAPS) ) @@ -16758,6 +16922,7 @@ async def test_team_info_reports_what_the_caller_may_edit(caller, org_admin, ena members_with_roles=[Member(user_id="admin-1", role="admin"), Member(user_id="member-1", role="user")], ) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) mock_prisma.get_data = AsyncMock(return_value=[]) @@ -16825,6 +16990,7 @@ _DB_OUTAGE_503_BODY: Final = { def _user_read_raising(error: Exception) -> tuple[MagicMock, MagicMock]: prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=error) cache = MagicMock() cache.async_get_cache = AsyncMock(return_value=None) diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py index 3da587435ad..d810f0bcdd7 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -107,6 +107,7 @@ class _FakePrisma: self.marker_landing_on_day: dict[str, str] = {} self.reconciled: list[str] = [] self.db = _FakeDb(self) + self.replica_db = self.db def write_late_row(self, day: str) -> None: """A per-key row for ``day`` lands now, after whatever scans already happened.""" diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index 89be341c87b..18c566906f6 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -73,6 +73,7 @@ def _query_raw_by_table( async def test_recover_double_hashed_key_metadata_via_active_token_digest(): double_hashed = hash_token("a" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _query_raw_by_table( active_rows=[_digest_row(double_hashed, "batch-worker", "team-1", "alice")], deleted_rows=[], @@ -91,6 +92,7 @@ async def test_recover_double_hashed_key_metadata_via_active_token_digest(): async def test_recover_double_hashed_key_metadata_falls_back_to_deleted_tokens(): double_hashed = hash_token("y" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _query_raw_by_table( active_rows=[], deleted_rows=[_digest_row(double_hashed, "deleted-key", "team-del", "erin")], @@ -109,6 +111,7 @@ async def test_recover_only_asks_deleted_tokens_for_digests_active_keys_missed() found_active = hash_token("1" * 64) found_deleted = hash_token("2" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _query_raw_by_table( active_rows=[_digest_row(found_active, "active-key", None, None)], deleted_rows=[_digest_row(found_deleted, "deleted-key", None, None)], @@ -128,6 +131,7 @@ async def test_recover_only_asks_deleted_tokens_for_digests_active_keys_missed() async def test_recover_permanent_miss_costs_two_digest_lookups_and_no_table_walk(): double_hashed = hash_token("b" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _query_raw_by_table(active_rows=[], deleted_rows=[]) result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}) @@ -141,6 +145,7 @@ async def test_recover_permanent_miss_costs_two_digest_lookups_and_no_table_walk @pytest.mark.asyncio async def test_recover_skips_keys_that_are_not_sha256_digests(): mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(return_value=[]) result = await recover_double_hashed_key_metadata(mock_prisma, {"sk-plain-key", "key-hash-short"}) @@ -153,6 +158,7 @@ async def test_recover_skips_keys_that_are_not_sha256_digests(): async def test_recover_returns_empty_when_digest_lookup_raises_prisma_error(): double_hashed = hash_token("c" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("db down")) result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}) @@ -164,6 +170,7 @@ async def test_recover_returns_empty_when_digest_lookup_raises_prisma_error(): async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows(): double_hashed = hash_token("d" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _query_raw_by_table( active_rows=[_digest_row(double_hashed, "recovered-alias", "team-9", "bob")], deleted_rows=[], @@ -202,6 +209,7 @@ async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows(): @pytest.mark.asyncio async def test_fill_missing_api_key_aliases_leaves_rows_untouched_when_nothing_is_missing(): mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(return_value=[]) rows = ({"api_key": hash_token("e" * 64), "api_key_alias": "named", "user_email": "x@example.com"},) @@ -215,6 +223,7 @@ async def test_fill_missing_api_key_aliases_leaves_rows_untouched_when_nothing_i async def test_fill_missing_api_key_aliases_keeps_spend_user_email_when_alias_is_missing(): double_hashed = hash_token("f" * 64) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = _query_raw_by_table( active_rows=[_digest_row(double_hashed, "team-key", "team-9", "key-owner")], deleted_rows=[], @@ -243,6 +252,7 @@ async def test_fill_missing_api_key_aliases_keeps_spend_user_email_when_alias_is @pytest.mark.asyncio async def test_fill_missing_api_key_aliases_skips_named_keys_that_have_no_email(): mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock(return_value=[]) rows = ( { @@ -264,6 +274,7 @@ async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_ session_digest = hash_token("cli-session-repro-user-6852") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction( mock_prisma, _query_raw_spend_logs( @@ -284,6 +295,7 @@ async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_ async def test_recover_key_metadata_from_spend_logs_skips_query_when_no_missing_keys(): window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[])) result = await recover_key_metadata_from_spend_logs(mock_prisma, set(), window, cache=InMemoryCache()) @@ -296,6 +308,7 @@ async def test_recover_key_metadata_from_spend_logs_skips_query_when_no_missing_ async def test_recover_key_metadata_from_spend_logs_returns_empty_on_prisma_error(): window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("db down"))) result = await recover_key_metadata_from_spend_logs( @@ -312,6 +325,7 @@ async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null foreign = hash_token("cli-session-foreign") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction( mock_prisma, _query_raw_spend_logs( @@ -334,6 +348,7 @@ async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null async def test_recover_key_metadata_from_spend_logs_skips_non_sha256_keys(): window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[])) result = await recover_key_metadata_from_spend_logs( @@ -349,6 +364,7 @@ async def test_recover_key_metadata_from_spend_logs_accepts_hashed_jwt_digests() jwt_digest = f"hashed-jwt-{hash_token('jwt-subject-1')}" window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(jwt_digest, None, "team-jwt", "jwt-user")])) result = await recover_key_metadata_from_spend_logs(mock_prisma, {jwt_digest}, window, cache=InMemoryCache()) @@ -366,6 +382,7 @@ async def test_recover_key_metadata_from_spend_logs_serves_repeat_lookups_from_t window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) cache = InMemoryCache() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, "owner-1")])) first = await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache) @@ -384,6 +401,7 @@ async def test_recover_key_metadata_from_spend_logs_only_queries_digests_the_cac window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) cache = InMemoryCache() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(cached_digest, "cached-alias", None, None)])) await recover_key_metadata_from_spend_logs(mock_prisma, {cached_digest}, window, cache=cache) query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(new_digest, "new-alias", None, None)])) @@ -401,6 +419,7 @@ async def test_recover_key_metadata_from_spend_logs_rescans_when_the_window_chan digest = hash_token("cli-session-windowed") cache = InMemoryCache() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([])) await recover_key_metadata_from_spend_logs( mock_prisma, {digest}, (datetime(2026, 9, 1), datetime(2026, 9, 4)), cache=cache @@ -421,6 +440,7 @@ async def test_recover_key_metadata_from_spend_logs_retries_a_failed_query_only_ window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("statement timeout"))) started = time.time() assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {} @@ -442,6 +462,7 @@ async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_ shared_ui_digest = hash_token("ui-token") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction( mock_prisma, _query_raw_spend_logs( @@ -461,6 +482,7 @@ async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_when_every_n digest = hash_token("cli-session-one-owner") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction( mock_prisma, _query_raw_spend_logs( @@ -480,6 +502,7 @@ async def test_recover_key_metadata_from_spend_logs_forgets_a_miss_long_before_a window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, None)])) started = time.time() @@ -498,6 +521,7 @@ async def test_recover_key_metadata_from_spend_logs_runs_one_query_for_concurren cache = InMemoryCache() lock = asyncio.Lock() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db async def slow_query_raw(sql: str, *params: object) -> list[dict[str, str | None]]: await asyncio.sleep(0.01) @@ -522,6 +546,7 @@ async def test_recover_key_metadata_from_spend_logs_keeps_a_repeated_miss_as_lon window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([])) await recover_key_metadata_from_spend_logs(mock_prisma, {unknown}, window, cache=cache) first_miss_key = next(key for key in cache.ttl_dict if unknown in key and not key.endswith(":missed-before")) @@ -539,6 +564,7 @@ async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_older_rows_a digest = hash_token("cli-session-owner-from-older-rows") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, None, "team-x", "alice")])) result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache()) @@ -551,6 +577,7 @@ async def test_recover_key_metadata_from_spend_logs_names_nothing_for_a_field_wh digest = hash_token("cli-session-disagreeing-rows") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db _spend_log_transaction( mock_prisma, _query_raw_spend_logs( @@ -576,6 +603,7 @@ async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_state digest = hash_token("cli-session-bounded-scan") window = (datetime(2026, 9, 7), datetime(2026, 9, 10)) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db calls: list[str] = [] transaction = MagicMock() transaction.execute_raw = AsyncMock(side_effect=lambda sql: calls.append(sql) or 0) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index c6a4173b583..b0b8ae671fc 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -245,6 +245,7 @@ def make_ui_spend_logs_mock_prisma(mock_spend_logs, filter_fn, team_lookup_fn=No class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db self.db.litellm_spendlogs = self.db if team_lookup_fn is not None: self.db.litellm_teamtable = self @@ -311,6 +312,7 @@ async def test_can_team_member_view_log_none_team_id(): def __init__(self): self.db = self.DB() + self.replica_db = self.db prisma = MockPrisma() auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1") @@ -334,6 +336,7 @@ async def test_can_team_member_view_log_team_not_found(monkeypatch): def __init__(self): self.db = self.DB() + self.replica_db = self.db prisma = MockPrisma() # Even if admin check would return True, no team means False @@ -375,6 +378,7 @@ async def test_can_team_member_view_log_not_admin(monkeypatch): def __init__(self): self.db = self.DB() + self.replica_db = self.db prisma = MockPrisma() monkeypatch.setattr( @@ -415,6 +419,7 @@ async def test_can_team_member_view_log_admin(monkeypatch): def __init__(self): self.db = self.DB() + self.replica_db = self.db prisma = MockPrisma() auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="user_1") @@ -476,6 +481,7 @@ def _make_owner_lookup_prisma(rows): class MockPrisma: def __init__(self): self.db = MockDB() + self.replica_db = self.db return MockPrisma() @@ -995,6 +1001,7 @@ async def test_ui_view_spend_logs_sort_by_and_sort_order( class MockPrismaClient: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_spendlogs = MagicMock() self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) self.db.query_raw = AsyncMock(side_effect=mock_query_raw) @@ -1059,6 +1066,7 @@ async def test_ui_view_spend_logs_sort_validation_errors( class MockPrismaClient: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_spendlogs = MagicMock() self.db.litellm_spendlogs.find_many = AsyncMock(return_value=[]) self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) @@ -1137,6 +1145,7 @@ async def test_ui_view_spend_logs_sort_by_request_duration_ms(client, monkeypatc class MockPrismaClient: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_spendlogs = MagicMock() self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) self.db.query_raw = AsyncMock(side_effect=mock_query_raw) @@ -1237,6 +1246,7 @@ async def test_ui_view_spend_logs_sort_by_model( class MockPrismaClient: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_spendlogs = MagicMock() self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) self.db.query_raw = AsyncMock(side_effect=mock_query_raw) @@ -1352,6 +1362,7 @@ async def test_ui_view_spend_logs_sort_by_ttft_ms(client, monkeypatch): class MockPrismaClient: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_spendlogs = MagicMock() self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) self.db.query_raw = AsyncMock(side_effect=mock_query_raw) @@ -2093,6 +2104,7 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db self.db.litellm_spendlogs = self.db mock_prisma_client = MockPrismaClient() @@ -2145,6 +2157,7 @@ async def test_ui_view_session_spend_logs_rehydrates_metadata_jsonb_text(client, class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db self.db.litellm_spendlogs = self.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) @@ -2186,6 +2199,7 @@ async def test_ui_view_session_spend_logs_scopes_non_admin_to_own_logs(client, m class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db self.db.litellm_spendlogs = self.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) @@ -2248,6 +2262,7 @@ async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, m class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db self.db.litellm_spendlogs = self.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) @@ -2736,6 +2751,7 @@ def _make_payload_lookup_prisma(rows): class MockPrisma: def __init__(self): self.db = MockDB() + self.replica_db = self.db return MockPrisma() @@ -2807,6 +2823,7 @@ async def test_ui_view_request_response_rejects_foreign_row_inserted_after_owner class MockPrisma: def __init__(self): self.db = MockDB() + self.replica_db = self.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrisma()) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( @@ -2838,6 +2855,7 @@ async def test_ui_view_request_response_custom_logger_denies_foreign_payload_own class MockPrisma: def __init__(self): self.db = MockDB() + self.replica_db = self.db class ColdStorageLogger: async def get_request_response_payload(self, request_id, start_time_utc, end_time_utc): @@ -3963,10 +3981,12 @@ async def test_global_spend_keys_endpoint_limit_validation(client, monkeypatch): """ # Create a simple mock for prisma client with empty response mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = MagicMock() mock_query_raw = AsyncMock(return_value=[]) mock_db.query_raw = mock_query_raw mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db # Apply the mock to the prisma_client module monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -4095,6 +4115,7 @@ async def test_view_spend_logs_summarize_parameter(client, monkeypatch): class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db # Apply the monkeypatch mock_prisma_client = MockPrismaClient() @@ -4194,6 +4215,7 @@ async def test_view_spend_logs_bounds_row_count(client, monkeypatch): class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db def hash_token(self, token): return f"hashed-{token}" @@ -4267,6 +4289,7 @@ async def test_view_spend_tags(client, monkeypatch): # Mock the prisma client and get_spend_by_tags function mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Mock response data @@ -4447,6 +4470,7 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch): class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) @@ -4511,6 +4535,7 @@ async def test_view_spend_logs_summarize_groups_by_day_in_sql(client, monkeypatc class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db def hash_token(self, token): return "hashed::" + token @@ -4578,6 +4603,7 @@ async def test_view_spend_logs_summarize_empty_rows(client, monkeypatch): class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( @@ -4619,6 +4645,7 @@ async def test_view_spend_logs_summarize_unhashed_api_key_without_padding(client class MockPrismaClient: def __init__(self): self.db = MockDB() + self.replica_db = self.db mock_prisma_client = MockPrismaClient() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -4921,6 +4948,7 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts(): ] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_spendlogs.group_by = AsyncMock() mock_prisma.db.query_raw = AsyncMock( return_value=[ @@ -4987,6 +5015,7 @@ async def test_build_ui_spend_logs_response_caps_session_models(): over_limit_models = [f"model-{i:02d}" for i in range(_SESSION_MODELS_LIMIT + 1)] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( return_value=[ { @@ -5040,6 +5069,7 @@ async def test_build_ui_spend_logs_response_key_split_session_gets_per_key_aggre ] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( return_value=[ { @@ -5104,6 +5134,7 @@ async def test_build_ui_spend_logs_response_empty_api_key_keeps_session_aggregat ] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( return_value=[ { @@ -5161,6 +5192,7 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_spend(): ] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # The raw aggregate query returns the full session spend (0.01 + 0.02 + 0.03). mock_prisma.db.query_raw = AsyncMock( return_value=[ @@ -5241,6 +5273,7 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_tokens(): ] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( return_value=[ { @@ -5323,6 +5356,7 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_duration(): ] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( return_value=[ { @@ -5382,6 +5416,7 @@ async def test_build_ui_spend_logs_response_session_cache_hit_count(): ] mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.query_raw = AsyncMock( return_value=[ { @@ -5451,6 +5486,7 @@ async def test_can_team_member_view_log_with_spend_logs_permission(monkeypatch): def __init__(self): self.db = self.DB() + self.replica_db = self.db prisma = MockPrisma() auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="member_1") @@ -5489,6 +5525,7 @@ async def test_can_team_member_view_log_without_spend_logs_permission(monkeypatc def __init__(self): self.db = self.DB() + self.replica_db = self.db prisma = MockPrisma() auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="member_1") @@ -5659,6 +5696,7 @@ class _CaptureFilterDB: class _CapturePrismaClient: def __init__(self): self.db = _CaptureFilterDB() + self.replica_db = self.db def hash_token(self, token): return "hashed::" + token @@ -5839,6 +5877,7 @@ class _SpendScopeMockPrismaClient: self.litellm_verificationtoken = _VerificationTokenTable() self.db = _DB() + self.replica_db = self.db async def get_data(self, table_name=None, query_type=None, **kwargs): self.get_data_calls.append( @@ -6227,6 +6266,7 @@ async def test_ui_view_spend_logs_rehydrates_metadata_jsonb_text(client, monkeyp class MockPrismaClient: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_spendlogs = MagicMock() self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) self.db.query_raw = AsyncMock(side_effect=mock_query_raw) @@ -6313,6 +6353,7 @@ async def test_ui_view_spend_logs_metadata_invalid_json_falls_back_to_empty_dict class MockPrismaClient: def __init__(self): self.db = MagicMock() + self.replica_db = self.db self.db.litellm_spendlogs = MagicMock() self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) self.db.query_raw = AsyncMock(side_effect=mock_query_raw) @@ -6624,7 +6665,7 @@ def test_ui_view_request_response_reads_from_cold_storage(client, monkeypatch): async def _query_raw(_sql, *_args): return [placeholder_row] - fake_prisma = SimpleNamespace(db=SimpleNamespace(query_raw=_query_raw)) + fake_prisma = SimpleNamespace(db=SimpleNamespace(query_raw=_query_raw), replica_db=SimpleNamespace(query_raw=_query_raw)) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", fake_prisma) cold_logger = _FakeColdStorageLogger( @@ -6680,6 +6721,7 @@ _SCOPED_SPEND_REPORT_PATHS = [ def _spend_report_mock_prisma(query_raw_returns=None, team_rows=None, user_row=None): pc = MagicMock() + pc.replica_db = pc.db pc.db.query_raw = AsyncMock( return_value=query_raw_returns if query_raw_returns is not None else [] ) @@ -7237,6 +7279,7 @@ def _session_grouped_mock_prisma(session_page_rows, session_total, representativ 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) return mock_prisma @@ -7282,6 +7325,7 @@ def _session_grouped_paginating_prisma(sessions, counted_total=None): 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) return mock_prisma @@ -7658,6 +7702,7 @@ async def test_ui_view_spend_logs_search_returns_flat_rows_when_grouping_by_sess 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) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( @@ -7704,6 +7749,7 @@ def _fake_prisma_with_owned_spend_log(owner_user_id, messages_json, response_jso class _Prisma: def __init__(self): self.db = _DB() + self.replica_db = self.db return _Prisma() @@ -7810,7 +7856,7 @@ def test_ui_view_request_response_internal_user_missing_row_forbidden(client, mo from types import SimpleNamespace - fake_prisma = SimpleNamespace(db=_DB()) + fake_prisma = SimpleNamespace(db=_DB(), replica_db=_DB()) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", fake_prisma) custom_logger = _RecordingAdditionalLoggingUtils( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 62ff08230d7..46dcaec4d66 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -106,6 +106,7 @@ def test_cors_exposes_cache_key_header_to_browser_js(): def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch): mock_login_result = {"user_id": "test-user"} mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_authenticate_user = AsyncMock(return_value=mock_login_result) mock_create_ui_token_object = MagicMock(return_value={"user_id": "test-user"}) mock_jwt_encode = MagicMock(return_value="signed-token") @@ -234,6 +235,7 @@ def test_login_v2_returns_json_on_proxy_exception(monkeypatch): from litellm.proxy._types import ProxyErrorTypes, ProxyException mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_authenticate_user = AsyncMock( side_effect=ProxyException( message="Invalid credentials", @@ -269,6 +271,7 @@ def test_login_v2_returns_json_on_http_exception(monkeypatch): from fastapi import HTTPException mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_authenticate_user = AsyncMock(side_effect=HTTPException(status_code=401, detail="Unauthorized")) monkeypatch.setattr( @@ -294,6 +297,7 @@ def test_login_v2_returns_json_on_http_exception(monkeypatch): def test_login_v2_returns_json_on_unexpected_exception(monkeypatch): """Test that /v2/login returns JSON error when unexpected exception occurs""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_authenticate_user = AsyncMock(side_effect=ValueError("Unexpected error")) monkeypatch.setattr( @@ -338,6 +342,7 @@ def test_login_v2_returns_json_on_invalid_json_body(monkeypatch): def test_login_v3_rejected_without_control_plane_url(monkeypatch): """v3/login returns 404 when control_plane_url is not configured.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key") monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -355,6 +360,7 @@ def test_login_v3_rejected_without_control_plane_url(monkeypatch): def test_login_v3_returns_code(monkeypatch): """v3/login returns an opaque code, not the JWT directly.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr( "litellm.proxy.auth.login_utils.authenticate_user", AsyncMock(return_value={"user_id": "test-user"}), @@ -393,6 +399,7 @@ def test_login_v3_returns_code(monkeypatch): def test_login_v3_exchange_happy_path(monkeypatch): """Full flow: v3/login returns code, v3/login/exchange redeems it for JWT.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr( "litellm.proxy.auth.login_utils.authenticate_user", AsyncMock(return_value={"user_id": "test-user"}), @@ -441,6 +448,7 @@ def test_login_v3_exchange_sets_secure_cookie_behind_trusted_tls_terminating_pro """Regression: /v3/login/exchange's token cookie must be Secure behind a trusted TLS-terminating reverse proxy even though litellm only sees a plain-HTTP hop.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr( "litellm.proxy.auth.login_utils.authenticate_user", AsyncMock(return_value={"user_id": "test-user"}), @@ -485,6 +493,7 @@ def test_login_v3_exchange_sets_secure_cookie_behind_trusted_tls_terminating_pro def test_login_v3_exchange_single_use(monkeypatch): """Code can only be redeemed once.""" mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db monkeypatch.setattr( "litellm.proxy.auth.login_utils.authenticate_user", AsyncMock(return_value={"user_id": "test-user"}), @@ -557,6 +566,7 @@ def test_login_v3_returns_json_on_proxy_exception(monkeypatch): from litellm.proxy._types import ProxyErrorTypes, ProxyException mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_authenticate_user = AsyncMock( side_effect=ProxyException( message="Invalid credentials", @@ -847,6 +857,7 @@ async def test_initialize_scheduled_jobs_loads_credentials_only_through_add_depl from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() mock_proxy_logging.db_spend_update_writer = MagicMock() @@ -906,6 +917,7 @@ async def test_periodic_reload_job_scheduled_without_store_model_in_db(monkeypat from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() @@ -948,6 +960,7 @@ async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval( from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() mock_proxy_logging.db_spend_update_writer = MagicMock() @@ -997,6 +1010,7 @@ async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_inte from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() mock_proxy_logging.db_spend_update_writer = MagicMock() @@ -1044,6 +1058,7 @@ async def test_initialize_scheduled_jobs_hydrates_mcp_when_store_model_in_db_fal from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_proxy_logging = MagicMock(spec=ProxyLogging) mock_proxy_logging.slack_alerting_instance = MagicMock() mock_proxy_logging.db_spend_update_writer = MagicMock() @@ -1907,10 +1922,12 @@ async def test_get_all_team_models(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = MagicMock() mock_litellm_teamtable = MagicMock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db mock_db.litellm_teamtable = mock_litellm_teamtable # Make find_many async @@ -2214,6 +2231,7 @@ async def test_non_admin_all_models_returns_user_models_when_user_row_missing(): user_added_model = {"model_name": "my-model", "model_info": {"id": "user-model-1"}} prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) prisma_client.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=MagicMock(created_by="ghost-user")) @@ -2349,6 +2367,7 @@ async def test_apply_search_filter_scopes_byok_to_caller_teams(): } prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=2) prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[db_caller_row, db_other_row]) caller_user_row = MagicMock() @@ -2432,6 +2451,7 @@ async def test_apply_search_filter_bounds_db_fetch_by_page_and_cap(): ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=10_000) prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) @@ -2477,6 +2497,7 @@ async def test_apply_search_filter_honours_exact_model_name_in_db_query(): from litellm.proxy.proxy_server import _apply_search_filter_to_models prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=0) prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) proxy_config = MagicMock() @@ -2558,6 +2579,7 @@ async def test_filter_models_by_team_id_excludes_viewer_direct_access(): } prisma = MagicMock() + prisma.replica_db = prisma.db team_db = MagicMock() team_db.model_dump.return_value = { "team_id": "team-111", @@ -2607,6 +2629,7 @@ async def test_filter_models_by_team_id_rejects_non_member(): } prisma = MagicMock() + prisma.replica_db = prisma.db # Caller is in team-222 only user_row = MagicMock() user_row.teams = ["team-222"] @@ -2645,6 +2668,7 @@ async def test_filter_models_by_team_id_allows_team_member(): } prisma = MagicMock() + prisma.replica_db = prisma.db user_row = MagicMock() user_row.teams = ["team-111", "team-999"] prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) @@ -2767,6 +2791,7 @@ async def test_add_access_group_models_to_team_models(): mock_ag_row.access_model_names = ["claude-3", "gemini"] mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_ag_row]) result = await _add_access_group_models_to_team_models( @@ -2843,6 +2868,7 @@ async def test_add_access_group_models_multiple_teams_shared_group(): mock_extra_row.access_model_names = ["gemini"] mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_shared_row, mock_extra_row]) result = await _add_access_group_models_to_team_models( @@ -2883,6 +2909,7 @@ async def test_add_access_group_models_no_eligible_teams(): team.access_group_ids = None mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock() result = await _add_access_group_models_to_team_models( @@ -2924,9 +2951,11 @@ async def test_get_all_team_models_with_access_groups(): mock_ag_row.access_model_names = ["claude-3"] mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_db = MagicMock() mock_litellm_teamtable = MagicMock() mock_prisma_client.db = mock_db + mock_prisma_client.replica_db = mock_prisma_client.db mock_db.litellm_teamtable = mock_litellm_teamtable mock_litellm_teamtable.find_many = AsyncMock(return_value=[mock_team1]) mock_db.litellm_accessgrouptable = MagicMock() @@ -3281,6 +3310,7 @@ async def test_add_proxy_budget_to_db_backfills_budget_reset_at(monkeypatch: pyt litellm_proxy_budget_name = "litellm-proxy-budget" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_usertable.update_many = AsyncMock(return_value={"count": 1}) mock_generate_key_helper = AsyncMock( @@ -4087,6 +4117,7 @@ async def test_write_config_to_file(monkeypatch): # Mock prisma_client to not be None (so DB path is taken) mock_prisma_client = AsyncMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.insert_data = AsyncMock() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -4456,6 +4487,7 @@ class TestPriceDataReloadAPI: ): # Mock the database connection with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.upsert = AsyncMock( return_value=_reload_schedule_row({}, reload_revision=1) ) @@ -4507,6 +4539,7 @@ class TestPriceDataReloadAPI: ): # Mock the database connection with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) mock_prisma.db.litellm_config.upsert = AsyncMock( return_value=_reload_schedule_row({}, reload_revision=1) @@ -4521,6 +4554,7 @@ class TestPriceDataReloadAPI: def test_schedule_model_cost_map_reload_admin_access(self, client_with_auth): """Admin schedule write owns param_value only, so it can't clobber the job-owned run columns""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db # Mock database upsert mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1)) @@ -4567,6 +4601,7 @@ class TestPriceDataReloadAPI: def test_cancel_model_cost_map_reload_admin_access(self, client_with_auth): """Test that admin users can cancel periodic reload""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=1) mock_prisma.db.litellm_config.delete = AsyncMock(return_value=None) @@ -4604,6 +4639,7 @@ class TestPriceDataReloadAPI: proxy_server_module.proxy_config.model_cost_map_loaded_at = datetime(2030, 6, 1, tzinfo=timezone.utc) with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row( {"interval_hours": 6}, @@ -4637,6 +4673,7 @@ class TestPriceDataReloadAPI: def test_get_model_cost_map_reload_status_no_config(self, client_with_auth): """Test that status returns not scheduled when no config exists""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) response = client_with_auth.get("/schedule/model_cost_map_reload/status") @@ -4651,6 +4688,7 @@ class TestPriceDataReloadAPI: def test_get_model_cost_map_reload_status_no_interval(self, client_with_auth): """A row left behind by a manual reload (no interval) must not read as scheduled""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row( {"interval_hours": None}, @@ -4670,6 +4708,7 @@ class TestPriceDataReloadAPI: def test_get_model_cost_map_reload_status_before_first_run(self, client_with_auth): """Scheduled but never executed: no last_run_at means no next_run can be computed""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row({"interval_hours": 6}) ) @@ -4726,6 +4765,7 @@ class TestPriceDataReloadIntegration: ): # Mock the database connection with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.upsert = AsyncMock( return_value=_reload_schedule_row({}, reload_revision=1) ) @@ -4764,6 +4804,7 @@ class TestPriceDataReloadIntegration: proxy_config = ProxyConfig() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1)) mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None) mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None) @@ -4822,6 +4863,7 @@ class TestPriceDataReloadIntegration: proxy_config = ProxyConfig() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None) mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row( @@ -4862,6 +4904,7 @@ class TestPriceDataReloadIntegration: proxy_config = ProxyConfig() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row( {"interval_hours": 6}, @@ -4911,6 +4954,7 @@ class TestPriceDataReloadIntegration: pods = [ProxyConfig(), ProxyConfig(), ProxyConfig()] frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None) mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1)) for pod in pods: @@ -4956,6 +5000,7 @@ class TestPriceDataReloadIntegration: frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc) proxy_config.model_cost_map_loaded_at = frozen_now mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None) mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row({}, reload_revision=published_revision) @@ -4992,6 +5037,7 @@ class TestPriceDataReloadIntegration: proxy_config = ProxyConfig() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=_reload_schedule_row({"interval_hours": 24})) mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1)) mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None) @@ -5033,6 +5079,7 @@ class TestPriceDataReloadIntegration: frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc) proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(hours=9) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7) ) @@ -5078,6 +5125,7 @@ class TestPriceDataReloadIntegration: pod_data_loaded_at = frozen_now - timedelta(hours=9) proxy_config.model_cost_map_loaded_at = pod_data_loaded_at mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7) ) @@ -5120,6 +5168,7 @@ class TestPriceDataReloadIntegration: frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc) proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(hours=9) mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_config.find_unique = AsyncMock( return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7) ) @@ -5211,6 +5260,7 @@ class TestPriceDataReloadIntegration: patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now), ): + mock_prisma.replica_db = mock_prisma.db mock_get_map.return_value = ModelCostMapReloaded( model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}} ) @@ -5252,6 +5302,7 @@ class TestPriceDataReloadIntegration: litellm_config_cache.flush_cache() proxy_config = ProxyConfig() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db # Set up config with interval_hours=12 and force_reload=True to trigger reload mock_config = MagicMock() @@ -5300,6 +5351,7 @@ class TestPriceDataReloadIntegration: mock_reload.return_value = {"anthropic": {"beta_header": "test-value"}} with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db # Simulate existing config with a schedule mock_existing = MagicMock() mock_existing.param_value = {"interval_hours": 8, "force_reload": False} @@ -5394,6 +5446,7 @@ async def test_add_router_settings_from_db_config_merge_logic(): # Mock prisma client mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) # Call the method under test @@ -5450,6 +5503,7 @@ async def test_invalid_db_routing_groups_do_not_abort_other_router_settings(): ], } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) await ProxyConfig()._add_router_settings_from_db_config(llm_router=router, prisma_client=mock_prisma_client) @@ -5472,6 +5526,7 @@ async def test_valid_db_routing_groups_still_replace_router_groups(): "routing_groups": [{"group_name": "g2", "models": ["m2"], "routing_strategy": "least-busy"}], } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) await ProxyConfig()._add_router_settings_from_db_config(llm_router=router, prisma_client=mock_prisma_client) @@ -5515,6 +5570,7 @@ async def test_add_router_settings_from_db_config_empty_db_lists_do_not_clobber_ } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) proxy_config.router_settings.load_yaml(config_data["router_settings"]) @@ -5552,6 +5608,7 @@ async def test_add_router_settings_from_db_config_empty_db_list_still_clears_unc mock_db_config.param_value = {"fallbacks": [], "model_group_alias": {}} mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) proxy_config.router_settings.load_yaml(config_data["router_settings"]) @@ -5598,6 +5655,7 @@ async def test_add_router_settings_from_db_config_edge_cases(): # Test Case 3: DB returns None (no router_settings in DB) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) config_data = {"router_settings": {"routing_strategy": "usage-based"}} @@ -5691,6 +5749,7 @@ async def test_add_router_settings_shallow_merge_behavior(): } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) proxy_config.router_settings.load_yaml(config_data["router_settings"]) @@ -5728,6 +5787,7 @@ async def test_router_settings_reload_keeps_db_values_writable(tmp_path, monkeyp return db_row if param_name == "router_settings" else None mock_prisma_client: Final = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=db_row) mock_router: Final = MagicMock() monkeypatch.setattr(proxy_server_module, "get_config_param", read_config_row) @@ -6353,6 +6413,7 @@ async def test_init_sso_settings_in_db(): } mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables @@ -6397,6 +6458,7 @@ async def test_init_sso_settings_in_db_no_settings(): # Mock prisma client to return None (no SSO settings) mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) # Mock _decrypt_and_set_db_env_variables @@ -6423,6 +6485,7 @@ async def test_init_sso_settings_in_db_error_handling(): # Mock prisma client to raise an exception mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=Exception("Database connection error")) # The method should not raise an exception, it should log it instead @@ -6451,6 +6514,7 @@ async def test_init_sso_settings_in_db_empty_settings(): mock_sso_config.sso_settings = {} mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config) # Mock _decrypt_and_set_db_env_variables @@ -6491,6 +6555,7 @@ async def test_init_sso_settings_in_db_retries_on_transport_error(): return mock_sso_config mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=_flaky_find_unique) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 @@ -6517,6 +6582,7 @@ async def test_init_sso_settings_in_db_propagates_when_reconnect_fails(): proxy_config = ProxyConfig() mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=prisma.errors.ClientNotConnectedError()) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 @@ -6548,6 +6614,7 @@ async def test_init_hashicorp_vault_config_override_retries_on_transport_error() return None # No config in DB → function returns early after retry. mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_configoverrides.find_unique = AsyncMock(side_effect=_flaky_find_unique) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 @@ -7131,6 +7198,7 @@ class TestInvitationEndpoints: def test_invitation_endpoints_proxy_admin_success(self, client_with_auth, endpoint, payload, mock_return): """Proxy admin can successfully create and delete invitations.""" with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_invitationlink = MagicMock() if endpoint == "/invitation/new": mock_create = AsyncMock(return_value=mock_return) @@ -7169,6 +7237,7 @@ class TestInvitationEndpoints: app.dependency_overrides[user_api_key_auth] = lambda: mock_auth with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_invitationlink = MagicMock() # Avoid triggering async DB calls in _user_has_admin_privileges with patch( @@ -8344,6 +8413,7 @@ async def test_batch_cost_poller_is_confirmed_before_serving(monkeypatch): from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=None) mock_proxy_logging = MagicMock(spec=ProxyLogging) @@ -8384,6 +8454,7 @@ async def test_store_model_in_db_db_override_when_config_false(): from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Mock DB returning store_model_in_db=True in general_settings mock_db_record = MagicMock() @@ -8429,6 +8500,7 @@ async def test_store_model_in_db_db_check_skipped_when_already_true(monkeypatch) from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) mock_proxy_logging = MagicMock(spec=ProxyLogging) @@ -8471,6 +8543,7 @@ async def test_store_model_in_db_db_failure_graceful(monkeypatch): from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db # Simulate DB failure mock_prisma_client.db.litellm_config.find_first = AsyncMock(side_effect=Exception("DB connection error")) @@ -8710,6 +8783,7 @@ async def test_prepare_spend_counter_increment_reseeds_from_db_on_counter_miss() db_row = MagicMock() db_row.spend = 42.0 fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row) stale_cache = DualCache() @@ -8811,6 +8885,7 @@ async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed( return row fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(side_effect=slow_find_unique) pod_a = DualCache() @@ -8872,6 +8947,7 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): org_row.spend = 305.0 fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) fake_prisma.db.litellm_endusertable.find_unique = AsyncMock() fake_prisma.db.litellm_tagtable.find_unique = AsyncMock() @@ -8904,6 +8980,7 @@ async def test_reseed_spend_from_db_skips_window_variant_keys(): from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_verificationtoken.find_unique = AsyncMock() fake_prisma.db.litellm_teamtable.find_unique = AsyncMock() @@ -8924,6 +9001,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss(): counter_cache = DualCache() window_start = datetime.now(timezone.utc) - timedelta(hours=1) fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None) fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( return_value=[{"api_key": "key-window", "_sum": {"spend": 2.25}}] @@ -8998,6 +9076,7 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory(): db_row = MagicMock() db_row.spend = 42.0 fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row) import litellm.proxy.proxy_server as ps @@ -9069,6 +9148,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None) fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( return_value=[{"api_key": "key-window-stale-local", "_sum": {"spend": 2.25}}] @@ -9146,6 +9226,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed() counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None) fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( return_value=[{"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}}] @@ -9444,6 +9525,7 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing(): db_row = MagicMock() db_row.spend = 362.0 fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row) import litellm.proxy.proxy_server as ps @@ -9548,6 +9630,7 @@ async def test_get_current_spend_coalesces_concurrent_reseeds(): counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique) import litellm.proxy.proxy_server as ps @@ -9585,6 +9668,7 @@ async def test_get_current_spend_uses_db_zero_over_stale_fallback(): db_row = MagicMock() db_row.spend = 0.0 fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row) import litellm.proxy.proxy_server as ps @@ -9655,6 +9739,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query(): counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique) import litellm.proxy.proxy_server as ps @@ -9770,6 +9855,7 @@ async def test_reseed_warms_cache_even_on_zero_db_spend(): return row fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=find_unique) import litellm.proxy.proxy_server as ps @@ -9833,6 +9919,7 @@ class _FakeLitellmConfig: class _FakePrismaClient: def __init__(self, initial_rows=None): self.db = mock.MagicMock() + self.replica_db = self.db self.db.litellm_config = _FakeLitellmConfig(initial_rows=initial_rows) self.jsonify_object = lambda obj: obj @@ -10698,6 +10785,7 @@ async def test_get_current_spend_redis_clean_miss_skips_stale_in_memory(): db_row = MagicMock() db_row.spend = 500.0 fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row) import litellm.proxy.proxy_server as ps @@ -10734,6 +10822,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory(): counter_cache.redis_cache = fake_redis fake_prisma = MagicMock() + fake_prisma.replica_db = fake_prisma.db fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=MagicMock(spend=999.0)) import litellm.proxy.proxy_server as ps @@ -11292,6 +11381,7 @@ class TestDeleteDeploymentSync: proxy_config = ProxyConfig() mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=Exception("DB connection lost")) result = await proxy_config._get_models_from_db(prisma_client=mock_prisma) @@ -11383,9 +11473,11 @@ def test_get_config_list_includes_cancel_on_disconnect(monkeypatch): from litellm.proxy.proxy_server import app mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_config_table = MagicMock() mock_config_table.find_first = AsyncMock(return_value=None) mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN @@ -11415,9 +11507,11 @@ def test_get_config_list_includes_apply_user_budget_to_team_keys(monkeypatch): from litellm.proxy.proxy_server import app mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_config_table = MagicMock() mock_config_table.find_first = AsyncMock(return_value=None) mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN @@ -11437,9 +11531,11 @@ def test_get_config_list_includes_user_api_key_cache_max_size(monkeypatch): """The Admin UI General Settings table renders whatever /config/list returns, so the cache capacity has to be exposed there as an Integer to be editable.""" mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_config_table = MagicMock() mock_config_table.find_first = AsyncMock(return_value=None) mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(proxy_server_module, "prisma_client", mock_prisma) app.dependency_overrides[proxy_server_module.user_api_key_auth] = lambda: UserAPIKeyAuth( user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN @@ -11468,9 +11564,11 @@ def test_get_config_list_includes_budget_exceeded_throttle_percentage(monkeypatc from litellm.proxy.proxy_server import app mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_config_table = MagicMock() mock_config_table.find_first = AsyncMock(return_value=None) mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.15) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( @@ -11545,9 +11643,11 @@ def test_get_config_list_includes_anthropic_prompt_caching_fields(monkeypatch): from litellm.proxy.proxy_server import app mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_config_table = MagicMock() mock_config_table.find_first = AsyncMock(return_value=None) mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", "1h") @@ -11853,9 +11953,11 @@ def test_get_config_list_marks_untouched_prompt_caching_flag_as_not_set(monkeypa from litellm.proxy.proxy_server import app mock_prisma = MagicMock() + mock_prisma.replica_db = mock_prisma.db mock_config_table = MagicMock() mock_config_table.find_first = AsyncMock(return_value=None) mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( @@ -12131,6 +12233,7 @@ def _config_field_info_client(monkeypatch, user_role): mock_config_table.find_first = AsyncMock(return_value=db_record) mock_prisma = MagicMock() mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + mock_prisma.replica_db = mock_prisma.db monkeypatch.setattr(ps, "prisma_client", mock_prisma) settings = SettingsStore("general_settings") @@ -12188,6 +12291,7 @@ def _fake_prisma_with_config(existing_param_value): """MagicMock prisma whose litellm_config row returns existing_param_value and whose litellm_auditlog.create records the written audit row.""" fake = MagicMock() + fake.replica_db = fake.db config_row = MagicMock() config_row.param_value = existing_param_value fake.db.litellm_config.find_first = AsyncMock(return_value=config_row) @@ -13709,6 +13813,7 @@ async def test_team_window_spend_row_carries_the_request_start_time(): def _mock_startup_prisma_client(health_check_error=None, connect_error=None): client = MagicMock() + client.replica_db = client.db client.connect = AsyncMock(side_effect=connect_error) client.db.start_token_refresh_task = AsyncMock() client.check_view_exists = AsyncMock() @@ -13849,6 +13954,7 @@ async def _run_scheduled_background_jobs(): from litellm.proxy.utils import ProxyLogging mock_prisma_client = MagicMock() + mock_prisma_client.replica_db = mock_prisma_client.db mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) mock_proxy_logging = MagicMock(spec=ProxyLogging) @@ -14242,6 +14348,7 @@ async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeyp return served_callback().prompt_manager.get_prompt("greeting_sync").content prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db try: prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[db_row("Begin every reply with AHOY")]) await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client) @@ -14286,6 +14393,7 @@ async def test_init_prompts_in_db_syncs_remaining_rows_when_one_row_fails(monkey return row prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db try: prisma_client.db.litellm_prompttable.find_many = AsyncMock( return_value=[db_row("broken_sync", "does_not_exist"), db_row("healthy_sync", "dotprompt")] @@ -14338,6 +14446,7 @@ async def test_init_prompts_in_db_syncs_every_environment_sharing_a_versioned_id return callback.prompt_manager.get_prompt("greeting_env").content prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db try: prisma_client.db.litellm_prompttable.find_many = AsyncMock( return_value=[ @@ -14398,6 +14507,7 @@ async def test_init_prompts_in_db_unloads_rows_deleted_on_another_worker(monkeyp monkeypatch.setattr(litellm, "callbacks", []) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db try: prisma_client.db.litellm_prompttable.find_many = AsyncMock( return_value=[_prompt_db_row("greeting_del", _dotprompt_params("greeting_del"))] @@ -14435,6 +14545,7 @@ async def test_init_prompts_in_db_keeps_config_prompts_when_their_id_has_no_db_r ) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db try: IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(prompt=config_prompt) prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[]) @@ -14457,6 +14568,7 @@ async def test_init_prompts_in_db_keeps_the_in_memory_copy_when_a_row_fails_to_p monkeypatch.setattr(litellm, "callbacks", []) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db try: prisma_client.db.litellm_prompttable.find_many = AsyncMock( return_value=[_prompt_db_row("greeting_broken", _dotprompt_params("greeting_broken"))] @@ -14489,6 +14601,7 @@ async def test_init_prompts_in_db_keeps_a_prompt_created_while_the_sync_was_read monkeypatch.setattr(litellm, "callbacks", []) prisma_client = MagicMock() + prisma_client.replica_db = prisma_client.db try: async def create_prompt_behind_the_select() -> list: diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 0fc7295a717..1005e423a28 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -2433,3 +2433,107 @@ def test_handle_exception_on_proxy_logs_bug_report_only_for_unmapped_500(caplog) assert provider_result.code == internal_result.code == "500" assert ISSUE_URL_BASE in caplog.text assert ISSUE_URL_BASE not in internal_result.message + + +class _ReadOnlyReplicaError(RuntimeError): + pass + + +class _FakeUserTable: + def __init__(self, rows: list[dict[str, str]], *, read_only: bool) -> None: + self._rows = rows + self._read_only = read_only + + async def find_many(self) -> list[dict[str, str]]: + return list(self._rows) + + async def find_unique(self, where: dict[str, str]) -> dict[str, str] | None: + return next((row for row in self._rows if row["user_id"] == where["user_id"]), None) + + async def create(self, data: dict[str, str]) -> dict[str, str]: + if self._read_only: + raise _ReadOnlyReplicaError("cannot execute INSERT in a read-only transaction") + self._rows.append(data) # mutable-ok: in-memory fake table + return data + + +class _FakePrisma: + def __init__(self, label: str, rows: list[dict[str, str]], *, read_only: bool, reachable: bool = True) -> None: + self._label = label + self._reachable = reachable + self._connected = False + self.litellm_usertable = _FakeUserTable(rows, read_only=read_only) + + def is_connected(self) -> bool: + return self._connected + + async def connect(self, timeout: object = None) -> None: + if not self._reachable: + raise ConnectionError(f"{self._label} unreachable") + self._connected = True + + async def query_raw(self, sql: str, *params: object) -> list[dict[str, str]]: + return [{"served_by": self._label}] + + +def _replica_client(*, reader_reachable: bool = True) -> tuple[PrismaClient, list[dict[str, str]], list[dict[str, str]]]: + from litellm.proxy.db.prisma_client import PrismaWrapper + from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper + + writer_rows: Final[list[dict[str, str]]] = [{"user_id": "u1", "user_email": "writer@example.com"}] + reader_rows: Final[list[dict[str, str]]] = [{"user_id": "u1", "user_email": "reader@example.com"}] + writer: Final = PrismaWrapper(original_prisma=_FakePrisma("writer", writer_rows, read_only=False)) + reader: Final = PrismaWrapper( + original_prisma=_FakePrisma("reader", reader_rows, read_only=True, reachable=reader_reachable) + ) + client: Final = PrismaClient.__new__(PrismaClient) + client.db = RoutingPrismaWrapper(writer=writer, reader=reader) + return client, writer_rows, reader_rows + + +class TestPrismaClientReplicaDb: + @pytest.mark.asyncio + async def test_replica_db_reads_come_from_the_reader(self) -> None: + client, _writer_rows, reader_rows = _replica_client() + await client.db.connect() + + found: Final = await client.replica_db.litellm_usertable.find_unique(where={"user_id": "u1"}) + raw: Final = await client.replica_db.query_raw("SELECT 1") + + assert found == reader_rows[0] + assert raw == [{"served_by": "reader"}] + + @pytest.mark.asyncio + async def test_replica_db_writes_land_on_the_writer(self) -> None: + client, writer_rows, reader_rows = _replica_client() + await client.db.connect() + new_row: Final = {"user_id": "u2", "user_email": "new@example.com"} + + created: Final = await client.replica_db.litellm_usertable.create(data=new_row) + + assert created == new_row + assert new_row in writer_rows + assert new_row not in reader_rows + + @pytest.mark.asyncio + async def test_replica_db_reads_fall_back_to_the_writer_when_the_reader_is_unreachable(self) -> None: + client, writer_rows, _reader_rows = _replica_client(reader_reachable=False) + await client.db.connect() + + found: Final = await client.replica_db.litellm_usertable.find_many() + raw: Final = await client.replica_db.query_raw("SELECT 1") + + assert found == writer_rows + assert raw == [{"served_by": "writer"}] + + @pytest.mark.asyncio + async def test_replica_db_is_the_plain_handle_without_a_replica(self) -> None: + from litellm.proxy.db.prisma_client import PrismaWrapper + + rows: Final[list[dict[str, str]]] = [{"user_id": "u1", "user_email": "only@example.com"}] + client: Final = PrismaClient.__new__(PrismaClient) + client.db = PrismaWrapper(original_prisma=_FakePrisma("writer", rows, read_only=False)) + + assert client.replica_db is client.db + assert await client.replica_db.litellm_usertable.find_many() == rows + assert await client.replica_db.query_raw("SELECT 1") == [{"served_by": "writer"}] diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py b/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py index 901fa8f57ff..c15f610e0bc 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_session_handler.py @@ -466,6 +466,7 @@ class _FakePrismaDB: class _FakePrismaClient: def __init__(self, results): self.db = _FakePrismaDB(results) + self.replica_db = self.db def _spend_log(request_id: str, session_id: str, prompt: str, answer: str) -> dict: