diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index a305d5be1e6..2775f4865a9 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -25,6 +25,8 @@ from typing import ( overload, ) +import prisma.errors + import litellm from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache, RedisCache @@ -64,6 +66,13 @@ else: PrismaClient = Any ProxyLogging = Any +PRISMA_DEADLOCK_CODE = "P2034" + + +def _is_deadlock_error(e: Exception) -> bool: + """Check if a Prisma error is a PostgreSQL deadlock / write conflict (P2034).""" + return isinstance(e, prisma.errors.DataError) and getattr(e, "code", None) == PRISMA_DEADLOCK_CODE + class DBSpendUpdateWriter: """ @@ -1084,10 +1093,7 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - user_id, - response_cost, - ) in user_list_transactions.items(): + for user_id, response_cost in sorted(user_list_transactions.items()): batcher.litellm_usertable.update_many( where={"user_id": user_id}, data={"spend": {"increment": response_cost}}, @@ -1102,9 +1108,12 @@ class DBSpendUpdateWriter: start_time=start_time, proxy_logging_obj=proxy_logging_obj, ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff + # Randomized backoff to reduce repeated collisions across pods + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + continue _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -1139,10 +1148,7 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - token, - response_cost, - ) in key_list_transactions.items(): + for token, response_cost in sorted(key_list_transactions.items()): batcher.litellm_verificationtoken.update_many( # 'update_many' prevents error from being raised if no row exists where={"token": token}, data={ @@ -1160,9 +1166,12 @@ class DBSpendUpdateWriter: start_time=start_time, proxy_logging_obj=proxy_logging_obj, ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff + # Randomized backoff to reduce repeated collisions across pods + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + continue _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -1183,10 +1192,7 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - team_id, - response_cost, - ) in team_list_transactions.items(): + for team_id, response_cost in sorted(team_list_transactions.items()): verbose_proxy_logger.debug( "Updating spend for team id={} by {}".format( team_id, response_cost @@ -1206,9 +1212,12 @@ class DBSpendUpdateWriter: start_time=start_time, proxy_logging_obj=proxy_logging_obj, ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff + # Randomized backoff to reduce repeated collisions across pods + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + continue _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -1241,10 +1250,7 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - key, - response_cost, - ) in team_member_list_transactions.items(): + for key, response_cost in sorted(team_member_list_transactions.items()): # key is "team_id::::user_id::" team_id = key.split("::")[1] user_id = key.split("::")[3] @@ -1264,9 +1270,12 @@ class DBSpendUpdateWriter: start_time=start_time, proxy_logging_obj=proxy_logging_obj, ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff + # Randomized backoff to reduce repeated collisions across pods + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + continue _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -1298,10 +1307,7 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - org_id, - response_cost, - ) in org_list_transactions.items(): + for org_id, response_cost in sorted(org_list_transactions.items()): batcher.litellm_organizationtable.update_many( # 'update_many' prevents error from being raised if no row exists where={"organization_id": org_id}, data={"spend": {"increment": response_cost}}, @@ -1316,16 +1322,16 @@ class DBSpendUpdateWriter: start_time=start_time, proxy_logging_obj=proxy_logging_obj, ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep( - # Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are - # cancelled basically at the same time, so if they wait the same time they will also retry at the same time - # and thus they are more likely to deadlock again. - # Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of - # repeated deadlocks, and therefore of exceeding the retry limit. - random.uniform(2**i, 2 ** (i + 1)) - ) + # Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are + # cancelled basically at the same time, so if they wait the same time they will also retry at the same time + # and thus they are more likely to deadlock again. + # Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of + # repeated deadlocks, and therefore of exceeding the retry limit. + await asyncio.sleep(random.uniform(2**i, 2 ** (i + 1))) except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2 ** (i + 1))) + continue _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -1389,7 +1395,7 @@ class DBSpendUpdateWriter: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for entity_id, response_cost in transactions.items(): + for entity_id, response_cost in sorted(transactions.items()): verbose_proxy_logger.debug( f"Updating spend for {entity_name} {where_field}={entity_id} by {response_cost}" ) @@ -1405,8 +1411,12 @@ class DBSpendUpdateWriter: start_time=start_time, proxy_logging_obj=proxy_logging_obj, ) - await asyncio.sleep(2**i) # Exponential backoff + # Randomized backoff to reduce repeated collisions across pods + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + continue _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -1709,14 +1719,17 @@ class DBSpendUpdateWriter: start_time=start_time, proxy_logging_obj=proxy_logging_obj, ) - await asyncio.sleep( - # Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are - # cancelled basically at the same time, so if they wait the same time they will also retry at the same time - # and thus they are more likely to deadlock again. - # Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of - # repeated deadlocks, and therefore of exceeding the retry limit. - random.uniform(2**i, 2 ** (i + 1)) - ) + # Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are + # cancelled basically at the same time, so if they wait the same time they will also retry at the same time + # and thus they are more likely to deadlock again. + # Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of + # repeated deadlocks, and therefore of exceeding the retry limit. + await asyncio.sleep(random.uniform(2**i, 2 ** (i + 1))) + except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2 ** (i + 1))) + continue + raise except Exception as e: if "transactions_to_process" in locals(): diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 01a0f55aac7..d324daefdee 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4,6 +4,7 @@ import hashlib import inspect import json import os +import random import smtplib import sys import threading @@ -24,6 +25,8 @@ from typing import ( overload, ) +import prisma.errors + from litellm import _custom_logger_compatible_callbacks_literal from litellm.constants import DEFAULT_MODEL_CREATED_AT_TIME, MAX_TEAM_LIST_LIMIT from litellm.proxy._types import ( @@ -34,9 +37,17 @@ from litellm.proxy._types import ( SpendLogsMetadata, SpendLogsPayload, ) + from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import CallTypes, CallTypesLiteral +PRISMA_DEADLOCK_CODE = "P2034" + + +def _is_deadlock_error(e: Exception) -> bool: + """Check if a Prisma error is a PostgreSQL deadlock / write conflict (P2034).""" + return isinstance(e, prisma.errors.DataError) and getattr(e, "code", None) == PRISMA_DEADLOCK_CODE + try: from litellm_enterprise.enterprise_callbacks.send_emails.base_email import ( BaseEmailLogger, @@ -4538,10 +4549,7 @@ class ProxyUpdateSpend: timeout=timedelta(seconds=60) ) as transaction: async with transaction.batch_() as batcher: - for ( - end_user_id, - response_cost, - ) in end_user_list_transactions.items(): + for end_user_id, response_cost in sorted(end_user_list_transactions.items()): if litellm.max_end_user_budget is not None: pass batcher.litellm_endusertable.upsert( @@ -4562,9 +4570,12 @@ class ProxyUpdateSpend: _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff + # Randomized backoff to reduce repeated collisions across pods + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) except Exception as e: + if _is_deadlock_error(e) and i < n_retry_times: + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + continue _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -4657,7 +4668,23 @@ class ProxyUpdateSpend: ) if i >= n_retry_times: raise - await asyncio.sleep(2**i) + # Randomized backoff to reduce repeated collisions across pods + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + except Exception as e: + if i is None: + i = 0 + if _is_deadlock_error(e) and i < n_retry_times: + verbose_proxy_logger.warning( + "Spend tracking - deadlock writing spend logs, " + "retry %d/%d. logs_count=%d, error=%s", + i + 1, + n_retry_times, + len(logs_to_process), + str(e), + ) + await asyncio.sleep(random.uniform(2**i, 2**(i+1))) + continue + raise except Exception as e: # Logs already removed from queue at start - don't put them back # This matches the original behavior where logs are removed even on error diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index b5f82ef04c5..2276ab274a6 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1218,6 +1218,137 @@ async def test_commit_key_spend_updates_includes_last_active(): assert before_call <= last_active <= after_call +@pytest.mark.asyncio +async def test_deadlock_error_retried_on_user_spend_update(): + """ + Test that a PostgreSQL deadlock (Prisma P2034) is retried with backoff + instead of immediately failing and losing spend data. + """ + import prisma.errors + + db_writer = DBSpendUpdateWriter() + + # First call raises a deadlock error, second call succeeds + deadlock_error = prisma.errors.DataError( + data={"user_facing_error": {"error_code": "P2034"}}, + message="Transaction failed due to a write conflict or a deadlock", + ) + + mock_batcher = MagicMock() + mock_batcher.litellm_usertable = MagicMock() + mock_batcher.litellm_usertable.update_many = MagicMock() + + mock_transaction_ok = AsyncMock() + mock_transaction_ok.__aenter__ = AsyncMock(return_value=mock_transaction_ok) + mock_transaction_ok.__aexit__ = AsyncMock(return_value=False) + mock_transaction_ok.batch_ = MagicMock( + return_value=AsyncMock( + __aenter__=AsyncMock(return_value=mock_batcher), + __aexit__=AsyncMock(return_value=False), + ) + ) + + mock_prisma_client = MagicMock() + # First call raises deadlock, second succeeds + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.tx = MagicMock( + side_effect=[ + # First attempt: raise deadlock inside the transaction context + AsyncMock( + __aenter__=AsyncMock(side_effect=deadlock_error), + __aexit__=AsyncMock(return_value=False), + ), + # Second attempt: succeed + mock_transaction_ok, + ] + ) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.call_details = {} + + db_spend_update_transactions = { + "user_list_transactions": {"user_1": 0.05}, + "end_user_list_transactions": {}, + "key_list_transactions": {}, + "team_list_transactions": {}, + "team_member_list_transactions": {}, + "org_list_transactions": {}, + "tag_list_transactions": {}, + "agent_list_transactions": {}, + } + + with patch("litellm.proxy.utils._raise_failed_update_spend_exception"), \ + patch("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", new_callable=AsyncMock) as mock_sleep: + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=3, + proxy_logging_obj=mock_proxy_logging, + db_spend_update_transactions=db_spend_update_transactions, + ) + + # Verify: tx was called twice (first deadlock, then success) + assert mock_prisma_client.db.tx.call_count == 2 + # Verify: sleep was called for backoff between retries + mock_sleep.assert_called() + + +@pytest.mark.asyncio +async def test_non_deadlock_prisma_error_not_retried(): + """ + Test that a non-deadlock Prisma error (e.g. UniqueViolationError) + is NOT retried — it should fail immediately. + """ + import prisma.errors + + db_writer = DBSpendUpdateWriter() + + # A non-deadlock DataError (e.g. unique violation, code P2002) + non_deadlock_error = prisma.errors.DataError( + data={"user_facing_error": {"error_code": "P2002"}}, + message="Unique constraint failed on the fields: (`user_id`)", + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.tx = MagicMock( + return_value=AsyncMock( + __aenter__=AsyncMock(side_effect=non_deadlock_error), + __aexit__=AsyncMock(return_value=False), + ) + ) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.call_details = {} + + db_spend_update_transactions = { + "user_list_transactions": {"user_1": 0.05}, + "end_user_list_transactions": {}, + "key_list_transactions": {}, + "team_list_transactions": {}, + "team_member_list_transactions": {}, + "org_list_transactions": {}, + "tag_list_transactions": {}, + "agent_list_transactions": {}, + } + + mock_raise = MagicMock() + with patch("litellm.proxy.utils._raise_failed_update_spend_exception", mock_raise), \ + patch("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", new_callable=AsyncMock): + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=3, + proxy_logging_obj=mock_proxy_logging, + db_spend_update_transactions=db_spend_update_transactions, + ) + + # Verify: _raise_failed_update_spend_exception was called immediately (no retry) + # The first call should be for the user table non-deadlock error + mock_raise.assert_called() + first_call_error = mock_raise.call_args_list[0][1]["e"] + assert isinstance(first_call_error, prisma.errors.DataError) + assert first_call_error.code == "P2002" + + @pytest.mark.asyncio async def test_update_database_creates_single_task(): """