mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): prevent DB deadlocks in concurrent spend updates
Sort entity IDs before batching to ensure consistent lock ordering across pods, detect PostgreSQL deadlocks (P2034) with targeted retry, and use randomized backoff to avoid repeated collisions.
This commit is contained in:
parent
3dccdde9c8
commit
49dcb169b2
3 changed files with 225 additions and 54 deletions
|
|
@ -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::<value>::user_id::<value>"
|
||||
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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue