mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(proxy): bound daily spend rollup row-lock waits with lock_timeout and requeue 55P03 (#44450)
Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
d181bc7b80
commit
2477635213
18 changed files with 704 additions and 222 deletions
|
|
@ -1807,6 +1807,7 @@ SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL",
|
|||
SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7))
|
||||
SPEND_LOG_WRITE_BATCH_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_BYTES", 2_000_000)))
|
||||
SPEND_LOG_WRITE_BATCH_MAX_ROWS: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_ROWS", "100")))
|
||||
SPEND_ROLLUP_LOCK_TIMEOUT_MS: Final = max(1, int(os.getenv("SPEND_ROLLUP_LOCK_TIMEOUT_MS", "5000")))
|
||||
SPEND_LOG_QUEUE_SIZE_THRESHOLD: Final = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
|
||||
SPEND_LOG_QUEUE_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_QUEUE_MAX_BYTES", "64000000")))
|
||||
SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0))
|
||||
|
|
|
|||
|
|
@ -75,6 +75,7 @@ from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
|
|||
)
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.db.model_usage_rollup import build_model_usage_transaction
|
||||
from litellm.proxy.db.rollup_lock_timeout import apply_rollup_lock_timeout
|
||||
from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING
|
||||
from litellm.proxy.spend_tracking.compression_savings import (
|
||||
extract_compression_saved_tokens,
|
||||
|
|
@ -292,6 +293,7 @@ async def _spend_update_tx(
|
|||
) -> AsyncGenerator[_SpendTransaction]:
|
||||
tx: Final[_SpendTransactionManager] = prisma_client.db.tx(timeout=timedelta(seconds=60))
|
||||
async with db_span(call_type, table), tx as transaction:
|
||||
await apply_rollup_lock_timeout(transaction)
|
||||
yield transaction
|
||||
|
||||
|
||||
|
|
@ -2032,11 +2034,17 @@ class DBSpendUpdateWriter:
|
|||
start_time: float,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> None:
|
||||
"""Retry a failed spend-update transaction on connection errors or deadlocks, else re-raise."""
|
||||
"""Retry a failed spend-update transaction on connection errors, deadlocks or a rollup
|
||||
``lock_timeout`` (55P03), else re-raise. All three roll the transaction back before any
|
||||
increment applied, so re-sending the same batch cannot double-count."""
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.utils import _raise_failed_update_spend_exception
|
||||
|
||||
is_retryable = isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) or PrismaDBExceptionHandler.is_deadlock_error(e)
|
||||
is_retryable = (
|
||||
isinstance(e, DB_RETRY_SAFE_ERROR_TYPES)
|
||||
or PrismaDBExceptionHandler.is_deadlock_error(e)
|
||||
or PrismaDBExceptionHandler.is_lock_timeout_error(e)
|
||||
)
|
||||
if not is_retryable or attempt >= n_retry_times:
|
||||
_raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
|
||||
verbose_proxy_logger.warning(
|
||||
|
|
|
|||
|
|
@ -31,6 +31,8 @@ _CONNECTION_CAPACITY_PHRASES: Final = (
|
|||
"remaining connection slots are reserved",
|
||||
)
|
||||
_PRISMA_POOL_TIMEOUT_CODE: Final = "P2024"
|
||||
_LOCK_TIMEOUT_SQLSTATE: Final = "55P03"
|
||||
_LOCK_TIMEOUT_PHRASE: Final = "canceling statement due to lock timeout"
|
||||
|
||||
|
||||
def _exception_chain(e: BaseException) -> Iterator[BaseException]:
|
||||
|
|
@ -266,6 +268,21 @@ class PrismaDBExceptionHandler:
|
|||
error_message: Final = str(e).lower()
|
||||
return any(phrase in error_message for phrase in _CONNECTION_CAPACITY_PHRASES)
|
||||
|
||||
@staticmethod
|
||||
def is_lock_timeout_error(e: Exception) -> bool:
|
||||
"""True iff Postgres cancelled a statement for exceeding ``lock_timeout``
|
||||
(SQLSTATE 55P03). The statement never acquired its lock, so it never
|
||||
applied and the transaction rolled back: the rows it carried are safe to
|
||||
re-send, and the pooled connection is free again rather than pinned
|
||||
behind the holder."""
|
||||
import prisma
|
||||
|
||||
if not isinstance(e, _exception_types(prisma.errors.PrismaError)):
|
||||
return False
|
||||
if PrismaDBExceptionHandler.postgres_sqlstate(e) == _LOCK_TIMEOUT_SQLSTATE:
|
||||
return True
|
||||
return _LOCK_TIMEOUT_PHRASE in str(e).lower()
|
||||
|
||||
@staticmethod
|
||||
def postgres_sqlstate(e: Exception) -> str | None:
|
||||
"""The SQLSTATE Postgres attached to a failed statement, as prisma surfaces it, or None."""
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from itertools import groupby
|
|||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import LiteralString
|
||||
|
||||
from litellm.constants import (
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
|
|
@ -17,6 +18,7 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES, SpendLogsPayload
|
||||
from litellm.proxy.db.model_insights_tasks import load_model_insight_tasks
|
||||
from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
|
@ -32,6 +34,8 @@ class _UpsertTable(Protocol):
|
|||
class _ModelUsageBatch(Protocol):
|
||||
litellm_dailymodelusage: _UpsertTable
|
||||
|
||||
def execute_raw(self, query: LiteralString, *args: object) -> None: ...
|
||||
|
||||
|
||||
class _ModelUsageBatchManager(Protocol):
|
||||
async def __aenter__(self) -> _ModelUsageBatch: ...
|
||||
|
|
@ -132,6 +136,7 @@ async def flush_model_usage_transactions(
|
|||
for attempt in range(n_retry_times + 1):
|
||||
try:
|
||||
async with _model_usage_batch(prisma_client) as batcher:
|
||||
batcher.execute_raw(ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting())
|
||||
for key, grouped in groupby(ordered, key=lambda transaction: transaction.key):
|
||||
entries = tuple(grouped)
|
||||
spend = sum(entry.spend for entry in entries)
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ _ROOT_FIELD: Final = re.compile(r"result:\s*(\w+)")
|
|||
_RAW_SQL: Final = re.compile(r'query:\s*"((?:[^"\\]|\\.)*)"')
|
||||
_LEADING_KEYWORD: Final = re.compile(r"(?:\\[nrt]|\s|\()*(\w+)")
|
||||
_SETTING: Final = re.compile(r"(?:\\[nrt]|\s)*SET\s+(?:LOCAL\s+|SESSION\s+)?([A-Za-z_.]+)", re.IGNORECASE)
|
||||
_SET_CONFIG: Final = re.compile(r"(?:\\[nrt]|\s)*SELECT\s+set_config\(\s*'([A-Za-z_.]+)'", re.IGNORECASE)
|
||||
_CATALOG: Final = re.compile(r"\bpg_\w+|\bto_regclass\b|\binformation_schema\b|\bcurrent_setting\s*\(|^\s*SHOW\b")
|
||||
_PROBE: Final = re.compile(r"(?:\\[nrt]|\s)*SELECT\s+\d+\s*;?(?:\\[nrt]|\s)*$", re.IGNORECASE)
|
||||
_CTE_WRITE: Final = re.compile(r"\b(UPDATE|INSERT|DELETE)\s+(?:INTO\s+|FROM\s+)?(?:\\?\")", re.IGNORECASE)
|
||||
|
|
@ -90,9 +91,13 @@ def sql_relation(sql: str) -> str | None:
|
|||
|
||||
def sql_operation(sql: str) -> tuple[str | None, str | None]:
|
||||
"""``(verb, target)`` for a raw statement: the SQL verb from its leading keyword and the
|
||||
relation it names, or for ``SET`` the setting it changes."""
|
||||
relation it names, or for ``SET`` (and its parameterizable twin ``SELECT set_config``)
|
||||
the setting it changes."""
|
||||
if _PROBE.match(sql):
|
||||
return "ping", None
|
||||
set_config: Final = _SET_CONFIG.match(sql)
|
||||
if set_config is not None:
|
||||
return "set", set_config.group(1).lower()
|
||||
keyword: Final = _LEADING_KEYWORD.match(sql)
|
||||
leading: Final = keyword.group(1).upper() if keyword is not None else ""
|
||||
cte_write: Final = _CTE_WRITE.search(sql) if leading == "WITH" else None
|
||||
|
|
|
|||
34
litellm/proxy/db/rollup_lock_timeout.py
Normal file
34
litellm/proxy/db/rollup_lock_timeout.py
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
"""The row-lock budget every spend rollup transaction runs under.
|
||||
|
||||
A rollup upsert that waits on a row another pod holds keeps its pooled
|
||||
connection for as long as it waits. ``lock_timeout`` bounds that wait, so a
|
||||
long holder costs the fleet one aborted statement per waiter instead of a
|
||||
connection pinned until the chain drains. The aborted statement never applied,
|
||||
so the caller requeues its rows (SQLSTATE 55P03, see
|
||||
``PrismaDBExceptionHandler.is_lock_timeout_error``).
|
||||
|
||||
``set_config(..., true)`` is ``SET LOCAL``: it lasts for the enclosing
|
||||
transaction only, and takes the value as a bind parameter, which ``SET``
|
||||
cannot. Issued as the first statement of every rollup transaction.
|
||||
"""
|
||||
|
||||
from typing import Final, Protocol
|
||||
|
||||
from typing_extensions import LiteralString
|
||||
|
||||
from litellm.constants import SPEND_ROLLUP_LOCK_TIMEOUT_MS
|
||||
|
||||
ROLLUP_LOCK_TIMEOUT_SQL: Final = "SELECT set_config('lock_timeout', $1::text, true)"
|
||||
|
||||
|
||||
def rollup_lock_timeout_setting() -> str:
|
||||
return f"{SPEND_ROLLUP_LOCK_TIMEOUT_MS}ms"
|
||||
|
||||
|
||||
class _RawExecutor(Protocol):
|
||||
async def execute_raw(self, query: LiteralString, *args: object) -> int: ...
|
||||
|
||||
|
||||
async def apply_rollup_lock_timeout(transaction: _RawExecutor) -> None:
|
||||
"""Bound the row-lock waits of the open ``transaction`` for the rest of its life."""
|
||||
_ = await transaction.execute_raw(ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting())
|
||||
|
|
@ -15,15 +15,18 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
import random
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from itertools import groupby
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
from typing_extensions import LiteralString
|
||||
|
||||
from litellm.constants import SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS
|
||||
from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES
|
||||
from litellm.proxy.db.db_span import db_span
|
||||
from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting
|
||||
from litellm.proxy.db.spend_log_batching import spend_log_write_batches
|
||||
from litellm.repositories.table_repositories import SpendLogToolIndexRepository
|
||||
|
||||
|
|
@ -99,6 +102,27 @@ def build_tool_usage_transaction(
|
|||
)
|
||||
|
||||
|
||||
class _ToolRollupTable(Protocol):
|
||||
def upsert(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ...
|
||||
|
||||
|
||||
class _ToolRollupBatch(Protocol):
|
||||
litellm_dailytoolspend: _ToolRollupTable
|
||||
|
||||
def execute_raw(self, query: LiteralString, *args: object) -> None: ...
|
||||
|
||||
|
||||
class _ToolRollupBatchManager(Protocol):
|
||||
async def __aenter__(self) -> _ToolRollupBatch: ...
|
||||
|
||||
async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ...
|
||||
|
||||
|
||||
def _tool_rollup_batch(prisma_client: PrismaClient) -> _ToolRollupBatchManager:
|
||||
batch: Final[_ToolRollupBatchManager] = prisma_client.db.batch_()
|
||||
return batch
|
||||
|
||||
|
||||
async def flush_tool_usage_transactions(
|
||||
prisma_client: PrismaClient,
|
||||
transactions: Sequence[ToolUsageTransaction],
|
||||
|
|
@ -140,8 +164,9 @@ async def flush_tool_usage_transactions(
|
|||
await index_table.create_many(data=statement_rows, skip_duplicates=True)
|
||||
async with (
|
||||
db_span("commit_daily_tool_spend", "LiteLLM_DailyToolSpend"),
|
||||
prisma_client.db.batch_() as batcher,
|
||||
_tool_rollup_batch(prisma_client) as batcher,
|
||||
):
|
||||
batcher.execute_raw(ROLLUP_LOCK_TIMEOUT_SQL, rollup_lock_timeout_setting())
|
||||
for (date_key, tool_name), grouped in groupby(per_tool_day, key=lambda entry: (entry[0], entry[1])):
|
||||
entries = tuple(grouped)
|
||||
spend = sum(entry[2] for entry in entries)
|
||||
|
|
|
|||
|
|
@ -7645,6 +7645,13 @@ async def update_daily_tag_spend(
|
|||
verbose_proxy_logger.error("Error updating daily tag spend: %s", e)
|
||||
|
||||
|
||||
def _rollup_batch_never_applied(e: Exception) -> bool:
|
||||
"""True when a drained rollup batch provably never reached its rows: Postgres had no
|
||||
connection for it (53300) or cancelled it under ``lock_timeout`` before it took the row
|
||||
lock (55P03). Either way re-sending the batch cannot double-count, so it is requeued."""
|
||||
return PrismaDBExceptionHandler.is_database_capacity_error(e) or PrismaDBExceptionHandler.is_lock_timeout_error(e)
|
||||
|
||||
|
||||
async def update_spend_logs_job(
|
||||
prisma_client: PrismaClient,
|
||||
db_writer_client: AsyncHTTPHandler | None,
|
||||
|
|
@ -7710,9 +7717,11 @@ async def _run_spend_logs_job(
|
|||
)
|
||||
|
||||
# Tool usage tracking: drain the request-time queue into the tool index and the
|
||||
# LiteLLM_DailyToolSpend rollup. A batch Postgres had no connection for was never
|
||||
# sent, so it is requeued and the job stops; any other failure is dropped because
|
||||
# a replay could double-count the rollup.
|
||||
# LiteLLM_DailyToolSpend rollup. A batch Postgres had no connection for, or whose
|
||||
# rollup statement was cancelled by lock_timeout, never applied, so it is requeued;
|
||||
# any other failure is dropped because a replay could double-count the rollup. Only
|
||||
# a full server (53300) stops the job: a timed-out row belongs to this rollup alone,
|
||||
# the writers below have their own rows and keep flushing.
|
||||
async with prisma_client._tool_usage_transactions_lock:
|
||||
tool_usage_to_process: Final = prisma_client.tool_usage_transactions[:MAX_LOGS_PER_INTERVAL]
|
||||
prisma_client.tool_usage_transactions = prisma_client.tool_usage_transactions[len(tool_usage_to_process) :]
|
||||
|
|
@ -7724,19 +7733,20 @@ async def _run_spend_logs_job(
|
|||
transactions=tool_usage_to_process,
|
||||
)
|
||||
except Exception as tool_tracking_err:
|
||||
if PrismaDBExceptionHandler.is_database_capacity_error(tool_tracking_err):
|
||||
if _rollup_batch_never_applied(tool_tracking_err):
|
||||
async with prisma_client._tool_usage_transactions_lock:
|
||||
prisma_client.tool_usage_transactions = [
|
||||
*tool_usage_to_process,
|
||||
*prisma_client.tool_usage_transactions,
|
||||
]
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend tracking - database out of connections during tool usage flush, "
|
||||
"requeued %s tool usage transactions for the next flush: %s",
|
||||
"Spend tracking - database out of connections or rollup row lock timed out during tool usage "
|
||||
"flush, requeued %s tool usage transactions for the next flush: %s",
|
||||
len(tool_usage_to_process),
|
||||
tool_tracking_err,
|
||||
)
|
||||
raise
|
||||
if PrismaDBExceptionHandler.is_database_capacity_error(tool_tracking_err):
|
||||
raise
|
||||
else:
|
||||
verbose_proxy_logger.error(
|
||||
"Spend tracking - tool usage flush failed; %s tool usage transactions dropped: %s",
|
||||
|
|
@ -7744,6 +7754,7 @@ async def _run_spend_logs_job(
|
|||
tool_tracking_err,
|
||||
)
|
||||
|
||||
# Model usage rollup: same requeue and stop rules as the tool rollup above
|
||||
async with prisma_client._model_usage_transactions_lock:
|
||||
model_usage_to_process: Final = prisma_client.model_usage_transactions
|
||||
prisma_client.model_usage_transactions = []
|
||||
|
|
@ -7752,11 +7763,28 @@ async def _run_spend_logs_job(
|
|||
|
||||
await flush_model_usage_transactions(prisma_client=prisma_client, transactions=model_usage_to_process)
|
||||
except Exception as model_usage_err:
|
||||
verbose_proxy_logger.error(
|
||||
"Spend tracking - model usage flush failed; %s model usage transactions dropped: %s",
|
||||
len(model_usage_to_process),
|
||||
model_usage_err,
|
||||
)
|
||||
if _rollup_batch_never_applied(model_usage_err):
|
||||
async with (
|
||||
prisma_client._model_usage_transactions_lock # pyright: ignore[reportPrivateUsage] # needed to requeue
|
||||
):
|
||||
prisma_client.model_usage_transactions = [
|
||||
*model_usage_to_process,
|
||||
*prisma_client.model_usage_transactions,
|
||||
]
|
||||
verbose_proxy_logger.warning(
|
||||
"Spend tracking - database out of connections or rollup row lock timed out during model usage "
|
||||
"flush, requeued %s model usage transactions for the next flush: %s",
|
||||
len(model_usage_to_process),
|
||||
model_usage_err,
|
||||
)
|
||||
if PrismaDBExceptionHandler.is_database_capacity_error(model_usage_err):
|
||||
raise
|
||||
else:
|
||||
verbose_proxy_logger.error(
|
||||
"Spend tracking - model usage flush failed; %s model usage transactions dropped: %s",
|
||||
len(model_usage_to_process),
|
||||
model_usage_err,
|
||||
)
|
||||
|
||||
await flush_baseline_accounting(prisma_client)
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.proxy.db.daily_spend_bulk_upsert import (
|
|||
merge_by_conflict_key,
|
||||
)
|
||||
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
||||
from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL
|
||||
|
||||
TAG_TABLE = DAILY_SPEND_TABLES["tag"]
|
||||
USER_TABLE = DAILY_SPEND_TABLES["user"]
|
||||
|
|
@ -146,8 +147,12 @@ def test_non_tag_tables_carry_no_request_id_column():
|
|||
class _RecordingDb:
|
||||
def __init__(self) -> None:
|
||||
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
||||
self.session_settings: list[tuple[str, tuple[object, ...]]] = []
|
||||
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
if query == ROLLUP_LOCK_TIMEOUT_SQL:
|
||||
self.session_settings.append((query, args))
|
||||
return 0
|
||||
self.statements.append((query, args))
|
||||
return len(args)
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from redis.exceptions import DataError
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceTypes
|
||||
from litellm.constants import SPEND_ROLLUP_LOCK_TIMEOUT_MS
|
||||
from litellm.proxy._types import DailyTagSpendTransaction, Litellm_EntityType, SpendUpdateQueueItem
|
||||
from litellm.proxy.db.db_spend_update_writer import (
|
||||
_TEAM_ADVISORY_LOCK_SQL,
|
||||
|
|
@ -33,6 +34,7 @@ from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdate
|
|||
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
|
||||
build_window_spend_transaction,
|
||||
)
|
||||
from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL
|
||||
from tests.unit.proxy.db.fake_prisma_engine import engine_call
|
||||
|
||||
|
||||
|
|
@ -103,11 +105,21 @@ async def test_update_database_attributes_router_rejected_failure_to_model_group
|
|||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.disable_spend_logs", True), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch("litellm.proxy.proxy_server.llm_router", llm_router), # test-quality-ok: get_llm_router reads this proxy_server module global at call time; no injection seam
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.disable_spend_logs", True
|
||||
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", MagicMock()
|
||||
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", MagicMock()
|
||||
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"
|
||||
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.llm_router", llm_router
|
||||
), # test-quality-ok: get_llm_router reads this proxy_server module global at call time; no injection seam
|
||||
):
|
||||
await db_writer.update_database(
|
||||
token="test-token",
|
||||
|
|
@ -320,7 +332,9 @@ async def test_a_routed_request_reaches_the_auto_router_rollup_whether_or_not_sp
|
|||
}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs
|
||||
), # test-quality-ok: update_database reads this proxy_server module global at call time; no injection seam
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma),
|
||||
patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"),
|
||||
patch(
|
||||
|
|
@ -356,9 +370,13 @@ class _RecordingDb:
|
|||
|
||||
def __init__(self, execute_raw: Callable[[], int] | None = None) -> None:
|
||||
self.statements: list[Statement] = []
|
||||
self.session_settings: list[Statement] = []
|
||||
self._execute_raw = execute_raw
|
||||
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
if query == ROLLUP_LOCK_TIMEOUT_SQL:
|
||||
self.session_settings.append((query, args))
|
||||
return 0
|
||||
self.statements.append((query, args))
|
||||
if self._execute_raw is not None:
|
||||
return self._execute_raw()
|
||||
|
|
@ -752,7 +770,7 @@ async def test_update_tag_db_with_valid_tags():
|
|||
"""
|
||||
Test that _update_tag_db correctly processes valid tags and adds them to the spend update queue.
|
||||
"""
|
||||
from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem
|
||||
from litellm.proxy._types import Litellm_EntityType
|
||||
|
||||
writer = DBSpendUpdateWriter()
|
||||
mock_prisma = MagicMock()
|
||||
|
|
@ -1048,7 +1066,7 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster
|
|||
),
|
||||
)
|
||||
|
||||
lock_call, spend_call = mock_transaction.execute_raw.await_args_list
|
||||
_lock_timeout_call, lock_call, spend_call = mock_transaction.execute_raw.await_args_list
|
||||
lock_statement, locked_team_id = lock_call.args
|
||||
assert lock_statement is _TEAM_ADVISORY_LOCK_SQL
|
||||
assert locked_team_id == team_id
|
||||
|
|
@ -1092,7 +1110,7 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u
|
|||
),
|
||||
)
|
||||
|
||||
*lock_calls, spend_call = mock_transaction.execute_raw.await_args_list
|
||||
_lock_timeout_call, *lock_calls, spend_call = mock_transaction.execute_raw.await_args_list
|
||||
_statement, members = spend_call.args
|
||||
assert [lock_call.args for lock_call in lock_calls] == [
|
||||
(_TEAM_ADVISORY_LOCK_SQL, "eng"),
|
||||
|
|
@ -1730,7 +1748,7 @@ async def test_endpoint_field_is_correctly_mapped_from_call_type():
|
|||
|
||||
for key, transaction in update_dict.items():
|
||||
# Verify endpoint is included in the key
|
||||
assert key == f"test-user_2024-01-01_test-key_gpt-4_openai_/chat/completions"
|
||||
assert key == "test-user_2024-01-01_test-key_gpt-4_openai_/chat/completions"
|
||||
|
||||
# Verify endpoint is set in the transaction
|
||||
assert transaction["endpoint"] == "/chat/completions"
|
||||
|
|
@ -2912,17 +2930,22 @@ async def test_daily_transaction_carries_compression_saved_tokens():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("estimate, recorded_savings, expected", [
|
||||
pytest.param(None, None, -0.005, id="plain-classifier-cost"),
|
||||
pytest.param({"version": 1, "status": "unknown"}, None, 0.0, id="unknown"),
|
||||
pytest.param({"version": 2, "status": "unknown"}, None, 0.0, id="unknown-v2"),
|
||||
pytest.param({"version": 1, "status": "unknown"}, -0.003, 0.0, id="unknown-stale-value"),
|
||||
pytest.param({"version": 0, "status": "estimated"}, -0.003, 0.0, id="unsupported-version"),
|
||||
pytest.param({"version": 1, "status": "estimated"}, -0.003, -0.003, id="estimated"),
|
||||
pytest.param(None, -0.003, -0.003, id="legacy"),
|
||||
])
|
||||
@pytest.mark.parametrize(
|
||||
"estimate, recorded_savings, expected",
|
||||
[
|
||||
pytest.param(None, None, -0.005, id="plain-classifier-cost"),
|
||||
pytest.param({"version": 1, "status": "unknown"}, None, 0.0, id="unknown"),
|
||||
pytest.param({"version": 2, "status": "unknown"}, None, 0.0, id="unknown-v2"),
|
||||
pytest.param({"version": 1, "status": "unknown"}, -0.003, 0.0, id="unknown-stale-value"),
|
||||
pytest.param({"version": 0, "status": "estimated"}, -0.003, 0.0, id="unsupported-version"),
|
||||
pytest.param({"version": 1, "status": "estimated"}, -0.003, -0.003, id="estimated"),
|
||||
pytest.param(None, -0.003, -0.003, id="legacy"),
|
||||
],
|
||||
)
|
||||
async def test_daily_transaction_compression_saved_tokens_zero_when_absent(
|
||||
estimate: dict[str, object] | None, recorded_savings: float | None, expected: float,
|
||||
estimate: dict[str, object] | None,
|
||||
recorded_savings: float | None,
|
||||
expected: float,
|
||||
) -> None:
|
||||
"""Requests without any compression metadata produce a zero count."""
|
||||
writer = DBSpendUpdateWriter()
|
||||
|
|
@ -2941,12 +2964,14 @@ async def test_daily_transaction_compression_saved_tokens_zero_when_absent(
|
|||
"prompt_tokens": 100,
|
||||
"completion_tokens": 10,
|
||||
"spend": 0.01,
|
||||
"metadata": json.dumps({
|
||||
"usage_object": {"prompt_tokens": 100, "completion_tokens": 10},
|
||||
"routing_decision": {"savings_baseline_model": "anthropic/claude-sonnet-5", "classifier_cost": 0.005},
|
||||
"autorouter_savings": recorded_savings,
|
||||
"autorouter_savings_estimate": estimate,
|
||||
}),
|
||||
"metadata": json.dumps(
|
||||
{
|
||||
"usage_object": {"prompt_tokens": 100, "completion_tokens": 10},
|
||||
"routing_decision": {"savings_baseline_model": "anthropic/claude-sonnet-5", "classifier_cost": 0.005},
|
||||
"autorouter_savings": recorded_savings,
|
||||
"autorouter_savings_estimate": estimate,
|
||||
}
|
||||
),
|
||||
}
|
||||
|
||||
transaction = await writer._common_add_spend_log_transaction_to_daily_transaction(
|
||||
|
|
@ -3417,6 +3442,9 @@ async def test_failed_per_entity_increment_from_redis_restores_only_what_may_sti
|
|||
def batch_(self):
|
||||
return _BatchContext()
|
||||
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
return 0
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
|
|
@ -3439,9 +3467,7 @@ async def test_failed_per_entity_increment_from_redis_restores_only_what_may_sti
|
|||
)
|
||||
|
||||
mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once()
|
||||
restored = mock_redis_update_buffer.restore_transactions_to_redis.call_args.kwargs[
|
||||
"db_spend_update_transactions"
|
||||
]
|
||||
restored = mock_redis_update_buffer.restore_transactions_to_redis.call_args.kwargs["db_spend_update_transactions"]
|
||||
assert restored["user_list_transactions"] is None
|
||||
assert restored["team_list_transactions"] == {"team-1": 1.5}
|
||||
assert restored["key_list_transactions"] == ({"key-1": 1.5} if safe_to_resend else None)
|
||||
|
|
@ -3985,6 +4011,77 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey
|
|||
proxy_logging.failure_handler.assert_not_called()
|
||||
|
||||
|
||||
_SPEND_UPDATE_PATHS: Final = [
|
||||
("key_list_transactions", "sk-abc"),
|
||||
("user_list_transactions", "user-1"),
|
||||
("team_list_transactions", "team-1"),
|
||||
("team_member_list_transactions", "team_id::team-1::user_id::user-1"),
|
||||
("org_list_transactions", "org-1"),
|
||||
("org_member_list_transactions", "organization_id::org-1::user_id::user-1"),
|
||||
("tag_list_transactions", "tag-1"),
|
||||
("agent_list_transactions", "agent-1"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("transactions_key, sample_key", _SPEND_UPDATE_PATHS)
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_spend_updates_retries_a_rollup_lock_timeout_then_commits(
|
||||
monkeypatch, transactions_key, sample_key
|
||||
):
|
||||
"""A spend UPDATE cancelled under ``lock_timeout`` (55P03) never took its row lock, so the
|
||||
transaction rolled back with nothing applied. Every entity path retries it like a deadlock
|
||||
and lands the increment exactly once, instead of raising and dropping the drained batch."""
|
||||
slept = []
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.db.db_spend_update_writer.asyncio.sleep",
|
||||
AsyncMock(side_effect=lambda s: slept.append(s)),
|
||||
)
|
||||
|
||||
mock_batcher = MagicMock()
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.tx = MagicMock(side_effect=[_failing_tx(_lock_timeout_error()), _good_tx(mock_batcher)])
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.failure_handler = AsyncMock()
|
||||
proxy_logging.call_details = {}
|
||||
|
||||
await DBSpendUpdateWriter()._commit_spend_updates_to_db(
|
||||
prisma_client=mock_prisma_client,
|
||||
n_retry_times=3,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
db_spend_update_transactions=_empty_spend_transactions(**{transactions_key: {sample_key: 0.5}}),
|
||||
)
|
||||
|
||||
assert {
|
||||
"transactions_opened": mock_prisma_client.db.tx.call_count,
|
||||
"backoff_sleeps": len(slept),
|
||||
"failure_handler_calls": proxy_logging.failure_handler.await_count,
|
||||
} == {"transactions_opened": 2, "backoff_sleeps": 1, "failure_handler_calls": 0}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_spend_updates_raises_after_exhausting_lock_timeout_retries(monkeypatch):
|
||||
"""A row that stays locked past every retry still surfaces as a failure rather than
|
||||
looping forever or being swallowed; the budget is the same one deadlocks get."""
|
||||
monkeypatch.setattr("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", AsyncMock(return_value=None))
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.tx = MagicMock(side_effect=lambda *a, **k: _failing_tx(_lock_timeout_error()))
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.failure_handler = AsyncMock()
|
||||
|
||||
with pytest.raises(PrismaDataError, match="lock timeout"):
|
||||
await DBSpendUpdateWriter()._commit_spend_updates_to_db(
|
||||
prisma_client=mock_prisma_client,
|
||||
n_retry_times=2,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
db_spend_update_transactions=_empty_spend_transactions(key_list_transactions={"sk-abc": 0.5}),
|
||||
)
|
||||
|
||||
assert mock_prisma_client.db.tx.call_count == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"call_type, expects_flush",
|
||||
|
|
@ -4881,3 +4978,92 @@ async def test_shutdown_drain_that_lands_before_the_interrupted_tag_commit_resol
|
|||
assert redis_buffer.restored == [drained], "a tag batch whose COMMIT came back failed must be restored to Redis"
|
||||
(upsert,) = _daily_upserts(final_db, "LiteLLM_DailyTagSpend")
|
||||
assert _row_values(upsert, "api_requests") == [1]
|
||||
|
||||
|
||||
def _lock_timeout_error() -> PrismaDataError:
|
||||
return PrismaDataError(
|
||||
data={
|
||||
"user_facing_error": {
|
||||
"is_panic": False,
|
||||
"message": "Error querying the database: canceling statement due to lock timeout",
|
||||
"meta": {"code": "55P03", "message": "canceling statement due to lock timeout"},
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_daily_spend_upsert_runs_under_the_rollup_lock_timeout():
|
||||
"""The bulk daily upsert opens its transaction by bounding how long it may wait on
|
||||
rows another pod holds; the setting is transaction-local and bound as a parameter."""
|
||||
prisma_client = _RecordingPrisma()
|
||||
daily_spend_transactions = {
|
||||
"key": {
|
||||
"user_id": "user-1",
|
||||
"date": "2024-01-01",
|
||||
"api_key": "test-api-key",
|
||||
"model": "gpt-4",
|
||||
"custom_llm_provider": "openai",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"spend": 0.1,
|
||||
"api_requests": 1,
|
||||
"successful_requests": 1,
|
||||
"failed_requests": 0,
|
||||
}
|
||||
}
|
||||
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
daily_spend_transactions=daily_spend_transactions,
|
||||
entity_type="user",
|
||||
entity_id_field="user_id",
|
||||
)
|
||||
|
||||
assert prisma_client.db.session_settings == [(ROLLUP_LOCK_TIMEOUT_SQL, (f"{SPEND_ROLLUP_LOCK_TIMEOUT_MS}ms",))]
|
||||
assert len(prisma_client.db.statements) == 1
|
||||
assert daily_spend_transactions == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_daily_spend_rows_survive_a_lock_timeout_for_the_next_flush():
|
||||
"""A statement cancelled under ``lock_timeout`` (55P03) never applied, so its rows
|
||||
are not data Postgres refused: they stay queued for the next flush instead of being
|
||||
dropped the way a constraint violation is."""
|
||||
|
||||
def cancelled_waiting_for_the_row() -> int:
|
||||
raise _lock_timeout_error()
|
||||
|
||||
prisma_client = _RecordingPrisma(execute_raw=cancelled_waiting_for_the_row)
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.failure_handler = AsyncMock()
|
||||
daily_spend_transactions = {
|
||||
"key": {
|
||||
"user_id": "user-1",
|
||||
"date": "2024-01-01",
|
||||
"api_key": "test-api-key",
|
||||
"model": "gpt-4",
|
||||
"custom_llm_provider": "openai",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"spend": 0.1,
|
||||
"api_requests": 1,
|
||||
"successful_requests": 1,
|
||||
"failed_requests": 0,
|
||||
}
|
||||
}
|
||||
|
||||
with pytest.raises(PrismaDataError, match="lock timeout"):
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
daily_spend_transactions=daily_spend_transactions,
|
||||
entity_type="user",
|
||||
entity_id_field="user_id",
|
||||
)
|
||||
|
||||
assert len(prisma_client.db.statements) == 1
|
||||
assert list(daily_spend_transactions) == ["key"]
|
||||
|
|
|
|||
|
|
@ -1,12 +1,10 @@
|
|||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request
|
||||
from prisma import errors as prisma_errors
|
||||
from prisma.engine.errors import BinaryNotFoundError, EngineConnectionError, EngineRequestError
|
||||
from prisma.errors import (
|
||||
|
|
@ -22,9 +20,7 @@ from prisma.errors import (
|
|||
UniqueViolationError,
|
||||
)
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
|
||||
|
|
@ -55,22 +51,12 @@ def test_is_database_infrastructure_error_prisma_connection_errors(prisma_error)
|
|||
PrismaError(),
|
||||
PrismaError("validation failed on query"),
|
||||
DataError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
UniqueViolationError(
|
||||
data={"user_facing_error": {"meta": {"table": "test_table"}}}
|
||||
),
|
||||
ForeignKeyViolationError(
|
||||
data={"user_facing_error": {"meta": {"table": "test_table"}}}
|
||||
),
|
||||
MissingRequiredValueError(
|
||||
data={"user_facing_error": {"meta": {"table": "test_table"}}}
|
||||
),
|
||||
UniqueViolationError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
ForeignKeyViolationError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
MissingRequiredValueError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
RawQueryError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
TableNotFoundError(
|
||||
data={"user_facing_error": {"meta": {"table": "test_table"}}}
|
||||
),
|
||||
RecordNotFoundError(
|
||||
data={"user_facing_error": {"meta": {"table": "test_table"}}}
|
||||
),
|
||||
TableNotFoundError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
RecordNotFoundError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
],
|
||||
)
|
||||
def test_is_database_transport_error_non_connection_prisma_errors(prisma_error):
|
||||
|
|
@ -82,12 +68,7 @@ def test_is_database_connection_generic_errors():
|
|||
"""
|
||||
Test non-Prisma error cases for database connection checking
|
||||
"""
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_connection_error(
|
||||
Exception("Regular error")
|
||||
)
|
||||
== False
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_database_connection_error(Exception("Regular error")) == False
|
||||
|
||||
# Test with ProxyException (DB connection)
|
||||
db_proxy_exception = ProxyException(
|
||||
|
|
@ -95,17 +76,11 @@ def test_is_database_connection_generic_errors():
|
|||
type=ProxyErrorTypes.no_db_connection,
|
||||
param="test-param",
|
||||
)
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_connection_error(db_proxy_exception)
|
||||
== True
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_database_connection_error(db_proxy_exception) == True
|
||||
|
||||
# Test with non-DB error
|
||||
regular_exception = Exception("Regular error")
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_connection_error(regular_exception)
|
||||
== False
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_database_connection_error(regular_exception) == False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -143,12 +118,7 @@ def test_is_database_service_unavailable_error_prisma_p1001_masquerades_as_datae
|
|||
}
|
||||
}
|
||||
)
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(
|
||||
p1001_as_dataerror
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(p1001_as_dataerror) is True
|
||||
|
||||
|
||||
def test_is_prisma_data_error_only_true_for_dataerror():
|
||||
|
|
@ -196,12 +166,7 @@ def test_is_database_service_unavailable_error_cached_plan_escapes_as_503():
|
|||
}
|
||||
}
|
||||
)
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(
|
||||
cached_plan_error
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(cached_plan_error) is True
|
||||
|
||||
|
||||
def test_is_database_service_unavailable_error_prisma_engine_malformed_payload():
|
||||
|
|
@ -229,10 +194,7 @@ def test_is_database_service_unavailable_error_prisma_engine_malformed_payload()
|
|||
prisma_engine_utils.handle_response_errors(None, malformed_payload)
|
||||
|
||||
assert "no attribute 'get'" in str(exc_info.value)
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value)
|
||||
is True
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) is True
|
||||
|
||||
|
||||
def test_is_prisma_engine_internal_error_excludes_application_attributeerror():
|
||||
|
|
@ -247,23 +209,15 @@ def test_is_prisma_engine_internal_error_excludes_application_attributeerror():
|
|||
with pytest.raises(AttributeError) as exc_info:
|
||||
application_bug()
|
||||
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_prisma_engine_internal_error(exc_info.value)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value)
|
||||
is False
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_prisma_engine_internal_error(exc_info.value) is False
|
||||
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(exc_info.value) is False
|
||||
|
||||
|
||||
def test_is_prisma_engine_internal_error_excludes_data_layer_prisma_error():
|
||||
"""A data-layer ``PrismaError`` (the DB IS reachable and rejected the data)
|
||||
must stay 401. These are always raised from prisma internals, so the check
|
||||
excludes any ``PrismaError`` by type before inspecting the traceback."""
|
||||
data_layer_error = UniqueViolationError(
|
||||
data={"user_facing_error": {"meta": {"table": "t"}}}
|
||||
)
|
||||
data_layer_error = UniqueViolationError(data={"user_facing_error": {"meta": {"table": "t"}}})
|
||||
with pytest.raises(UniqueViolationError) as exc_info:
|
||||
raise data_layer_error
|
||||
e = exc_info.value
|
||||
|
|
@ -284,9 +238,7 @@ def test_is_database_service_unavailable_error_excludes_non_infra(error):
|
|||
"""Data-layer errors (the DB IS reachable and answered) and generic
|
||||
non-DB errors must NOT be classified as service-unavailable, otherwise a
|
||||
genuine 401 would be masked as a transient 503."""
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is False
|
||||
)
|
||||
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(error) is False
|
||||
|
||||
|
||||
def _wrapped_like_get_user_object(original):
|
||||
|
|
@ -393,22 +345,14 @@ def test_is_database_service_unavailable_error_asyncpg(monkeypatch):
|
|||
monkeypatch.setitem(sys.modules, "asyncpg.exceptions", fake_exceptions)
|
||||
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(
|
||||
PostgresConnectionError("connection reset")
|
||||
)
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(PostgresConnectionError("connection reset"))
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(
|
||||
InterfaceError("connection was closed")
|
||||
)
|
||||
is True
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(InterfaceError("connection was closed")) is True
|
||||
)
|
||||
assert (
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(
|
||||
UniqueViolationError("duplicate key")
|
||||
)
|
||||
is False
|
||||
PrismaDBExceptionHandler.is_database_service_unavailable_error(UniqueViolationError("duplicate key")) is False
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -675,7 +619,10 @@ def test_is_deadlock_error_excludes_non_deadlocks(error):
|
|||
"22021",
|
||||
),
|
||||
(RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"message": "m"}}}), None),
|
||||
(RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"code": 42, "message": "m"}}}), None),
|
||||
(
|
||||
RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"code": 42, "message": "m"}}}),
|
||||
None,
|
||||
),
|
||||
(prisma_errors.DataError(data={"user_facing_error": {"meta": None}}), None),
|
||||
(
|
||||
prisma_errors.DataError(
|
||||
|
|
@ -695,7 +642,9 @@ def test_is_deadlock_error_excludes_non_deadlocks(error):
|
|||
(httpx.ReadTimeout("no reply"), None),
|
||||
],
|
||||
)
|
||||
def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement(error: Exception, sqlstate: str | None):
|
||||
def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement(
|
||||
error: Exception, sqlstate: str | None
|
||||
):
|
||||
"""Only a prisma data error carrying Postgres's own error code yields a SQLSTATE, whether in ``meta``
|
||||
or, for a batched statement, only in the message; a codeless or malformed payload, an engine-level
|
||||
error, and a transport error yield None."""
|
||||
|
|
@ -905,7 +854,7 @@ def _pool_timeout_error() -> DataError:
|
|||
RawQueryError(
|
||||
data={
|
||||
"user_facing_error": {
|
||||
"message": 'Raw query failed. Code: `53300`. Message: `db error: FATAL: sorry, too many clients already`',
|
||||
"message": "Raw query failed. Code: `53300`. Message: `db error: FATAL: sorry, too many clients already`",
|
||||
"meta": {"code": "53300", "message": "FATAL: sorry, too many clients already"},
|
||||
"error_code": "P2010",
|
||||
}
|
||||
|
|
@ -945,3 +894,60 @@ def test_capacity_error_is_seen_through_a_wrapping_exception() -> None:
|
|||
wrapped: Final = RuntimeError("spend flush failed")
|
||||
wrapped.__cause__ = _capacity_error("sorry, too many clients already")
|
||||
assert PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(wrapped) is True
|
||||
|
||||
|
||||
def _lock_timeout_error(
|
||||
message: str = "canceling statement due to lock timeout", code: str | None = "55P03"
|
||||
) -> DataError:
|
||||
"""The shape prisma raises for a statement Postgres cancelled under ``lock_timeout``."""
|
||||
user_facing: Final[dict[str, object]] = {"is_panic": False, "message": f"Error querying the database: {message}"}
|
||||
if code is not None:
|
||||
user_facing["meta"] = {"code": code, "message": message}
|
||||
return DataError(data={"user_facing_error": user_facing})
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
[
|
||||
_lock_timeout_error(),
|
||||
_lock_timeout_error(message="Abbruch der Anweisung wegen Zeitüberschreitung beim Warten auf eine Sperre"),
|
||||
_lock_timeout_error(code=None),
|
||||
DataError(
|
||||
data={
|
||||
"user_facing_error": {
|
||||
"message": 'Error occurred during query execution: ConnectorError(ConnectorError { user_facing_error: None, kind: QueryError(PostgresError { code: "55P03", message: "canceling statement due to lock timeout", severity: "ERROR" }) })'
|
||||
}
|
||||
}
|
||||
),
|
||||
],
|
||||
ids=["meta_55P03", "meta_55P03_localised_message", "message_only", "batched_statement"],
|
||||
)
|
||||
def test_is_lock_timeout_error_recognises_a_statement_cancelled_under_lock_timeout(error: DataError) -> None:
|
||||
"""SQLSTATE 55P03 means the statement never took its lock, so it never applied
|
||||
and its rows are safe to re-send. It is a wait budget, not a full server, so
|
||||
it stays out of the connection-capacity classification."""
|
||||
assert PrismaDBExceptionHandler.is_lock_timeout_error(error) is True
|
||||
assert PrismaDBExceptionHandler.is_database_capacity_error(error) is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
[
|
||||
_lock_timeout_error(message="canceling statement due to statement timeout", code="57014"),
|
||||
_lock_timeout_error(message="deadlock detected", code="40P01"),
|
||||
_capacity_error("sorry, too many clients already"),
|
||||
_pool_timeout_error(),
|
||||
httpx.ConnectError("connection refused"),
|
||||
RuntimeError("canceling statement due to lock timeout"),
|
||||
],
|
||||
ids=[
|
||||
"statement_timeout_57014",
|
||||
"deadlock_40P01",
|
||||
"capacity_53300",
|
||||
"pool_timeout_P2024",
|
||||
"transport",
|
||||
"not_prisma",
|
||||
],
|
||||
)
|
||||
def test_is_lock_timeout_error_excludes_other_failures(error: Exception) -> None:
|
||||
assert PrismaDBExceptionHandler.is_lock_timeout_error(error) is False
|
||||
|
|
|
|||
|
|
@ -64,6 +64,9 @@ class _FakeTransaction:
|
|||
def batch_(self) -> _FakeBatchManager:
|
||||
return _FakeBatchManager(self._batcher)
|
||||
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
return 0
|
||||
|
||||
async def __aenter__(self) -> "_FakeTransaction":
|
||||
return self
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from unittest.mock import MagicMock
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.db import rollup_lock_timeout
|
||||
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
||||
from litellm.proxy.db.model_usage_rollup import (
|
||||
ModelUsageKey,
|
||||
|
|
@ -14,11 +15,13 @@ from litellm.proxy.db.model_usage_rollup import (
|
|||
flush_model_usage_transactions,
|
||||
model_usage_task_type,
|
||||
)
|
||||
from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL
|
||||
|
||||
|
||||
class _FakeBatcher:
|
||||
def __init__(self) -> None:
|
||||
self.litellm_dailymodelusage = MagicMock()
|
||||
self.execute_raw = MagicMock()
|
||||
|
||||
async def __aenter__(self) -> "_FakeBatcher":
|
||||
return self
|
||||
|
|
@ -201,3 +204,26 @@ async def test_request_time_path_queues_usage_instead_of_writing_to_the_db() ->
|
|||
|
||||
assert [transaction.key.model for transaction in prisma.model_usage_transactions] == ["openai/gpt-5.4-mini"]
|
||||
prisma.db.litellm_dailymodelusage.upsert.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rollup_transaction_sets_the_lock_timeout_before_its_first_upsert(monkeypatch) -> None:
|
||||
"""The DailyModelUsage rollup shares the DailyToolSpend write shape, so it runs
|
||||
under the same per-transaction lock budget, issued before any upsert queues."""
|
||||
monkeypatch.setattr(rollup_lock_timeout, "SPEND_ROLLUP_LOCK_TIMEOUT_MS", 250)
|
||||
batcher = _FakeBatcher()
|
||||
prisma = _prisma(MagicMock(return_value=batcher))
|
||||
order: list[str] = []
|
||||
batcher.execute_raw.side_effect = lambda *_: order.append("lock_timeout")
|
||||
batcher.litellm_dailymodelusage.upsert.side_effect = lambda **_: order.append("upsert")
|
||||
|
||||
await flush_model_usage_transactions(
|
||||
prisma_client=prisma,
|
||||
transactions=[
|
||||
build_model_usage_transaction(_payload()),
|
||||
build_model_usage_transaction(_payload(model="other")),
|
||||
],
|
||||
)
|
||||
|
||||
assert batcher.execute_raw.call_args_list == [((ROLLUP_LOCK_TIMEOUT_SQL, "250ms"),)]
|
||||
assert order == ["lock_timeout", "upsert", "upsert"]
|
||||
|
|
|
|||
|
|
@ -106,6 +106,7 @@ def test_a_payload_the_parser_does_not_know_stays_the_legacy_function_named_span
|
|||
),
|
||||
('INSERT INTO "LiteLLM_DailyUserSpend" (id) VALUES ($1)', ("insert", "LiteLLM_DailyUserSpend")),
|
||||
("SET LOCAL lock_timeout = 1000", ("set", "lock_timeout")),
|
||||
("SELECT set_config('lock_timeout', $1::text, true)", ("set", "lock_timeout")),
|
||||
("SELECT COUNT(*) FROM pg_stat_activity", ("select", "pg_catalog")),
|
||||
("SELECT 1", ("ping", None)),
|
||||
("SELECT current_setting('transaction_read_only') AS transaction_read_only", ("select", "pg_catalog")),
|
||||
|
|
|
|||
|
|
@ -4,22 +4,23 @@ only) and the flush that writes LiteLLM_SpendLogToolIndex in bounded statements
|
|||
plus the LiteLLM_DailyToolSpend rollup in one transaction.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from collections.abc import Awaitable, Callable
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.db import spend_log_tool_index
|
||||
from litellm.proxy.db import rollup_lock_timeout, spend_log_tool_index
|
||||
from litellm.proxy.db.log_db_metrics import record_db_io
|
||||
from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL
|
||||
from litellm.proxy.db.spend_log_tool_index import (
|
||||
ToolUsageTransaction,
|
||||
build_tool_usage_transaction,
|
||||
flush_tool_usage_transactions,
|
||||
response_tool_call_names,
|
||||
)
|
||||
from litellm.proxy.db.log_db_metrics import record_db_io
|
||||
from tests.unit.proxy.db.fake_prisma_engine import engine_call
|
||||
|
||||
|
||||
|
|
@ -32,6 +33,7 @@ class _FakeBatcher:
|
|||
def __init__(self) -> None:
|
||||
self.litellm_spendlogtoolindex = MagicMock()
|
||||
self.litellm_dailytoolspend = MagicMock()
|
||||
self.execute_raw = MagicMock()
|
||||
|
||||
async def __aenter__(self) -> "_FakeBatcher":
|
||||
return self
|
||||
|
|
@ -84,7 +86,9 @@ class TestBuildToolUsageTransaction:
|
|||
mcp_namespaced_tool_name=None,
|
||||
spend=0.5,
|
||||
total_tokens=100,
|
||||
completion_response=SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(tool_calls=None))]),
|
||||
completion_response=SimpleNamespace(
|
||||
choices=[SimpleNamespace(message=SimpleNamespace(tool_calls=None))]
|
||||
),
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
|
@ -395,3 +399,25 @@ async def test_a_tool_usage_flush_renders_one_postgres_span_per_table_written(
|
|||
"postgres.insert LiteLLM_SpendLogToolIndex",
|
||||
"postgres.upsert LiteLLM_DailyToolSpend",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rollup_transaction_sets_the_lock_timeout_before_its_first_upsert(monkeypatch) -> None:
|
||||
"""Every DailyToolSpend rollup transaction opens by bounding how long its upserts
|
||||
may wait on a row another pod holds, so a waiter costs one cancelled statement
|
||||
rather than a pooled connection pinned until the holder commits. The budget is the
|
||||
SPEND_ROLLUP_LOCK_TIMEOUT_MS knob, bound as a parameter, not baked into the SQL."""
|
||||
monkeypatch.setattr(rollup_lock_timeout, "SPEND_ROLLUP_LOCK_TIMEOUT_MS", 250)
|
||||
prisma, batcher = _prisma_with_batcher()
|
||||
order: list[str] = []
|
||||
batcher.execute_raw.side_effect = lambda *_: order.append("lock_timeout")
|
||||
batcher.litellm_dailytoolspend.upsert.side_effect = lambda **_: order.append("upsert")
|
||||
|
||||
await flush_tool_usage_transactions(
|
||||
prisma_client=prisma,
|
||||
transactions=[_transaction("r1", tool_names=("tool_a", "tool_b"))],
|
||||
)
|
||||
|
||||
assert batcher.execute_raw.call_args_list == [((ROLLUP_LOCK_TIMEOUT_SQL, "250ms"),)]
|
||||
assert order == ["lock_timeout", "upsert", "upsert"]
|
||||
assert "set_config('lock_timeout'" in ROLLUP_LOCK_TIMEOUT_SQL and ", true)" in ROLLUP_LOCK_TIMEOUT_SQL
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
from typing import Dict
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.utils import update_daily_tag_spend
|
||||
from litellm.proxy._types import DailyTagSpendTransaction
|
||||
import httpx
|
||||
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
||||
from litellm.proxy.db.rollup_lock_timeout import ROLLUP_LOCK_TIMEOUT_SQL
|
||||
from litellm.proxy.utils import update_daily_tag_spend
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -18,9 +18,7 @@ async def test_update_daily_tag_spend_delegates_to_tag_commit_writer():
|
|||
proxy_logging_obj.db_spend_update_writer = MagicMock()
|
||||
proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock()
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = (
|
||||
AsyncMock()
|
||||
)
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock()
|
||||
|
||||
await update_daily_tag_spend(
|
||||
prisma_client,
|
||||
|
|
@ -43,12 +41,8 @@ async def test_update_daily_tag_spend_logs_error_and_does_not_raise():
|
|||
redis_update_buffer._should_commit_spend_updates_to_redis.return_value = False
|
||||
proxy_logging_obj.db_spend_update_writer = MagicMock()
|
||||
proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock(
|
||||
side_effect=ValueError("boom")
|
||||
)
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = (
|
||||
AsyncMock()
|
||||
)
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock(side_effect=ValueError("boom"))
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock()
|
||||
|
||||
with patch("litellm.proxy.utils.verbose_proxy_logger.error") as error_logger:
|
||||
await update_daily_tag_spend(
|
||||
|
|
@ -69,9 +63,7 @@ async def test_update_daily_tag_spend_uses_redis_writer_when_enabled():
|
|||
proxy_logging_obj.db_spend_update_writer = MagicMock()
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock()
|
||||
proxy_logging_obj.db_spend_update_writer.redis_update_buffer = redis_update_buffer
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = (
|
||||
AsyncMock()
|
||||
)
|
||||
proxy_logging_obj.db_spend_update_writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock()
|
||||
|
||||
await update_daily_tag_spend(
|
||||
prisma_client,
|
||||
|
|
@ -92,17 +84,20 @@ async def test_daily_tag_spend_retries_then_succeeds():
|
|||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Fail the upsert 3 times with retryable DB errors, then succeed.
|
||||
prisma_client.db.execute_raw = AsyncMock(
|
||||
side_effect=[
|
||||
httpx.ConnectError("x"),
|
||||
httpx.ConnectError("x"),
|
||||
httpx.ConnectError("x"),
|
||||
1,
|
||||
]
|
||||
)
|
||||
upsert_outcomes = iter([httpx.ConnectError("x"), httpx.ConnectError("x"), httpx.ConnectError("x"), 1])
|
||||
|
||||
async def execute_raw(query: str, *args: object) -> int:
|
||||
if query == ROLLUP_LOCK_TIMEOUT_SQL:
|
||||
return 0
|
||||
outcome = next(upsert_outcomes)
|
||||
if isinstance(outcome, Exception):
|
||||
raise outcome
|
||||
return outcome
|
||||
|
||||
prisma_client.db.execute_raw = AsyncMock(side_effect=execute_raw)
|
||||
prisma_client.db.tx.return_value.__aenter__.return_value.execute_raw = prisma_client.db.execute_raw
|
||||
|
||||
daily_spend_transactions: Dict[str, DailyTagSpendTransaction] = {
|
||||
daily_spend_transactions: dict[str, DailyTagSpendTransaction] = {
|
||||
"k": {
|
||||
"tag": "prod-tag",
|
||||
"date": "2026-04-03",
|
||||
|
|
@ -135,7 +130,8 @@ async def test_daily_tag_spend_retries_then_succeeds():
|
|||
daily_spend_transactions=daily_spend_transactions,
|
||||
)
|
||||
|
||||
assert prisma_client.db.execute_raw.await_count == 4
|
||||
upsert_attempts = [c for c in prisma_client.db.execute_raw.await_args_list if c.args[0] != ROLLUP_LOCK_TIMEOUT_SQL]
|
||||
assert len(upsert_attempts) == 4
|
||||
assert sleep_mock.await_count == 3
|
||||
# The batch is one statement, so the successful attempt is a single call carrying
|
||||
# the row rather than one call per key.
|
||||
|
|
|
|||
|
|
@ -225,6 +225,9 @@ class _InMemoryBatcher:
|
|||
def __init__(self, table: _InMemoryUsageTable) -> None:
|
||||
self.litellm_dailymodelusage = table
|
||||
|
||||
def execute_raw(self, query: str, *args: object) -> None:
|
||||
return None
|
||||
|
||||
async def __aenter__(self) -> "_InMemoryBatcher":
|
||||
return self
|
||||
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ import json
|
|||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, Final, List
|
||||
from typing import Any, Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -111,15 +111,11 @@ async def test_update_daily_tag_spend_redis_path_when_buffered(
|
|||
writer = MagicMock()
|
||||
proxy_logging.db_spend_update_writer = writer
|
||||
writer.redis_update_buffer = MagicMock()
|
||||
writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(
|
||||
return_value=True
|
||||
)
|
||||
writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=True)
|
||||
writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock()
|
||||
writer._commit_daily_tag_spend_to_db = AsyncMock()
|
||||
|
||||
await update_daily_tag_spend(
|
||||
prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging
|
||||
)
|
||||
await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging)
|
||||
redis_kwargs = writer._commit_daily_tag_spend_to_db_with_redis.await_args.kwargs
|
||||
pinned = {
|
||||
"redis_calls": writer._commit_daily_tag_spend_to_db_with_redis.await_count,
|
||||
|
|
@ -130,9 +126,7 @@ async def test_update_daily_tag_spend_redis_path_when_buffered(
|
|||
assert pinned == {
|
||||
"redis_calls": 1,
|
||||
"direct_calls": 0,
|
||||
"redis_kwargs_keys": sorted(
|
||||
["prisma_client", "n_retry_times", "proxy_logging_obj"]
|
||||
),
|
||||
"redis_kwargs_keys": sorted(["prisma_client", "n_retry_times", "proxy_logging_obj"]),
|
||||
"redis_n_retries": 3,
|
||||
}
|
||||
|
||||
|
|
@ -145,15 +139,11 @@ async def test_update_daily_tag_spend_direct_path_when_no_redis(
|
|||
writer = MagicMock()
|
||||
proxy_logging.db_spend_update_writer = writer
|
||||
writer.redis_update_buffer = MagicMock()
|
||||
writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(
|
||||
return_value=False
|
||||
)
|
||||
writer.redis_update_buffer._should_commit_spend_updates_to_redis = MagicMock(return_value=False)
|
||||
writer._commit_daily_tag_spend_to_db_with_redis = AsyncMock()
|
||||
writer._commit_daily_tag_spend_to_db = AsyncMock()
|
||||
|
||||
await update_daily_tag_spend(
|
||||
prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging
|
||||
)
|
||||
await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging)
|
||||
assert writer._commit_daily_tag_spend_to_db.await_count == 1
|
||||
assert writer._commit_daily_tag_spend_to_db_with_redis.await_count == 0
|
||||
|
||||
|
|
@ -175,9 +165,7 @@ async def test_update_daily_tag_spend_logs_and_swallows_errors(
|
|||
proxy_logging.db_spend_update_writer._commit_daily_tag_spend_to_db = AsyncMock(
|
||||
side_effect=RuntimeError("commit boom")
|
||||
)
|
||||
await update_daily_tag_spend(
|
||||
prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging
|
||||
)
|
||||
await update_daily_tag_spend(prisma_client=mock_prisma_client, proxy_logging_obj=proxy_logging)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -265,15 +253,11 @@ async def test_update_spend_logs_job_processes_and_clears_queue(
|
|||
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock()
|
||||
|
||||
# Stub auxiliary imports so the test focuses on the spend-logs write path.
|
||||
import litellm.proxy.guardrails.usage_tracking as guard_mod
|
||||
import litellm.proxy.db.spend_log_tool_index as tool_mod
|
||||
import litellm.proxy.guardrails.usage_tracking as guard_mod
|
||||
|
||||
monkeypatch.setattr(
|
||||
guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
tool_mod, "flush_tool_usage_transactions", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False)
|
||||
monkeypatch.setattr(tool_mod, "flush_tool_usage_transactions", AsyncMock(), raising=False)
|
||||
|
||||
await update_spend_logs_job(
|
||||
prisma_client=mock_prisma_client,
|
||||
|
|
@ -283,12 +267,10 @@ async def test_update_spend_logs_job_processes_and_clears_queue(
|
|||
pinned = {
|
||||
"create_many_calls": mock_prisma_client.db.litellm_spendlogs.create_many.await_count,
|
||||
"queue_after": mock_prisma_client.spend_log_transactions,
|
||||
"first_data_request_id": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[
|
||||
"data"
|
||||
][0]["request_id"],
|
||||
"skip_duplicates_set": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs[
|
||||
"skip_duplicates"
|
||||
"first_data_request_id": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs["data"][0][
|
||||
"request_id"
|
||||
],
|
||||
"skip_duplicates_set": mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs["skip_duplicates"],
|
||||
}
|
||||
assert pinned == {
|
||||
"create_many_calls": 1,
|
||||
|
|
@ -315,9 +297,7 @@ async def test_update_spend_logs_job_requeues_popped_rows_when_write_cancelled(
|
|||
mock_prisma_client.spend_log_transactions.append(row_arriving_mid_flush)
|
||||
raise asyncio.CancelledError()
|
||||
|
||||
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(
|
||||
side_effect=_cancel_mid_write
|
||||
)
|
||||
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_cancel_mid_write)
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await update_spend_logs_job(
|
||||
|
|
@ -326,9 +306,7 @@ async def test_update_spend_logs_job_requeues_popped_rows_when_write_cancelled(
|
|||
proxy_logging_obj=proxy_logging,
|
||||
)
|
||||
|
||||
assert [
|
||||
row["request_id"] for row in mock_prisma_client.spend_log_transactions
|
||||
] == ["r1", "r2", "r3"]
|
||||
assert [row["request_id"] for row in mock_prisma_client.spend_log_transactions] == ["r1", "r2", "r3"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -369,12 +347,8 @@ async def test_drain_spend_logs_queue_flushes_rows_queued_while_draining(
|
|||
import litellm.proxy.db.spend_log_tool_index as tool_mod
|
||||
import litellm.proxy.guardrails.usage_tracking as guard_mod
|
||||
|
||||
monkeypatch.setattr(
|
||||
guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False)
|
||||
monkeypatch.setattr(tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.failure_handler = AsyncMock()
|
||||
|
|
@ -385,9 +359,7 @@ async def test_drain_spend_logs_queue_flushes_rows_queued_while_draining(
|
|||
async def _write(*args: Any, **kwargs: Any) -> None:
|
||||
written.extend(row["request_id"] for row in kwargs["data"])
|
||||
if len(written) == 1:
|
||||
mock_prisma_client.spend_log_transactions.append(
|
||||
make_spend_log_row(request_id="r2")
|
||||
)
|
||||
mock_prisma_client.spend_log_transactions.append(make_spend_log_row(request_id="r2"))
|
||||
|
||||
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write)
|
||||
|
||||
|
|
@ -408,12 +380,8 @@ async def test_drain_spend_logs_queue_stops_monitor_and_keeps_its_popped_rows(
|
|||
import litellm.proxy.db.spend_log_tool_index as tool_mod
|
||||
import litellm.proxy.guardrails.usage_tracking as guard_mod
|
||||
|
||||
monkeypatch.setattr(
|
||||
guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False)
|
||||
monkeypatch.setattr(tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.failure_handler = AsyncMock()
|
||||
|
|
@ -460,12 +428,8 @@ async def test_drain_spend_logs_queue_gives_up_after_max_passes(
|
|||
import litellm.proxy.db.spend_log_tool_index as tool_mod
|
||||
import litellm.proxy.guardrails.usage_tracking as guard_mod
|
||||
|
||||
monkeypatch.setattr(
|
||||
guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False
|
||||
)
|
||||
monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False)
|
||||
monkeypatch.setattr(tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False)
|
||||
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.failure_handler = AsyncMock()
|
||||
|
|
@ -474,9 +438,7 @@ async def test_drain_spend_logs_queue_gives_up_after_max_passes(
|
|||
async def _write_and_refill(*args: Any, **kwargs: Any) -> None:
|
||||
mock_prisma_client.spend_log_transactions.append(make_spend_log_row())
|
||||
|
||||
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(
|
||||
side_effect=_write_and_refill
|
||||
)
|
||||
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write_and_refill)
|
||||
|
||||
await drain_spend_logs_queue(
|
||||
prisma_client=mock_prisma_client,
|
||||
|
|
@ -484,10 +446,7 @@ async def test_drain_spend_logs_queue_gives_up_after_max_passes(
|
|||
proxy_logging_obj=proxy_logging,
|
||||
)
|
||||
|
||||
assert (
|
||||
mock_prisma_client.db.litellm_spendlogs.create_many.await_count
|
||||
== MAX_SPEND_LOG_DRAIN_ITERATIONS
|
||||
)
|
||||
assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == MAX_SPEND_LOG_DRAIN_ITERATIONS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -496,8 +455,8 @@ async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty(
|
|||
make_spend_log_row: Any,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
import litellm.proxy.utils as utils_mod
|
||||
import litellm.constants as constants_mod
|
||||
import litellm.proxy.utils as utils_mod
|
||||
|
||||
monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False)
|
||||
monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_SIZE_THRESHOLD", 1, raising=False)
|
||||
|
|
@ -530,8 +489,8 @@ async def test_monitor_spend_logs_queue_swallows_errors_and_backs_off(
|
|||
"""An exception inside the loop is logged with backoff and the loop
|
||||
continues running rather than crashing the monitor task.
|
||||
"""
|
||||
import litellm.proxy.utils as utils_mod
|
||||
import litellm.constants as constants_mod
|
||||
import litellm.proxy.utils as utils_mod
|
||||
|
||||
monkeypatch.setattr(constants_mod, "SPEND_LOG_QUEUE_POLL_INTERVAL", 0.0, raising=False)
|
||||
|
||||
|
|
@ -721,8 +680,7 @@ def test_raise_failed_update_spend_exception_emits_failure_handler() -> None:
|
|||
else None
|
||||
),
|
||||
"non_blocking_in_traceback": (
|
||||
"Non-Blocking"
|
||||
in proxy_logging.failure_handler.call_args.kwargs["traceback_str"]
|
||||
"Non-Blocking" in proxy_logging.failure_handler.call_args.kwargs["traceback_str"]
|
||||
if proxy_logging.failure_handler.call_args
|
||||
else False
|
||||
),
|
||||
|
|
@ -1086,3 +1044,152 @@ async def test_tool_usage_flush_still_drops_ambiguous_failures(
|
|||
)
|
||||
|
||||
assert mock_prisma_client.tool_usage_transactions == []
|
||||
|
||||
|
||||
def _postgres_lock_timeout() -> Exception:
|
||||
"""prisma's shape for Postgres SQLSTATE 55P03: the rollup upsert waited past
|
||||
``lock_timeout`` for a row another pod held and was cancelled before it took the lock."""
|
||||
from prisma.errors import DataError
|
||||
|
||||
return DataError(
|
||||
data={
|
||||
"user_facing_error": {
|
||||
"is_panic": False,
|
||||
"message": "Error querying the database: canceling statement due to lock timeout",
|
||||
"meta": {"code": "55P03", "message": "canceling statement due to lock timeout"},
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_usage_flush_requeues_when_the_rollup_row_lock_times_out(
|
||||
mock_prisma_client: MagicMock,
|
||||
) -> None:
|
||||
"""A rollup transaction cancelled by ``lock_timeout`` never applied: the waiter
|
||||
gives its connection back instead of pinning it behind the holder, and its
|
||||
increments go back to the queue head for the next flush rather than being dropped.
|
||||
Unlike a full server, one contended row says nothing about the other writers, so the
|
||||
job carries on: the model usage rollup still opens its own transaction and the
|
||||
auto-router drain still runs."""
|
||||
mock_prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock()
|
||||
rollup_batch = MagicMock()
|
||||
rollup_batch.__aenter__ = AsyncMock(return_value=MagicMock())
|
||||
rollup_batch.__aexit__ = AsyncMock(side_effect=_postgres_lock_timeout())
|
||||
mock_prisma_client.db.batch_ = MagicMock(return_value=rollup_batch)
|
||||
mock_prisma_client.spend_log_transactions = []
|
||||
first, second = _tool_usage_transaction("r1"), _tool_usage_transaction("r2")
|
||||
mock_prisma_client.tool_usage_transactions = [first, second]
|
||||
mock_prisma_client.model_usage_transactions = [_model_usage_transaction("gpt-4o")]
|
||||
mock_prisma_client.autorouter_turn_transactions = [MagicMock(name="autorouter_turn")]
|
||||
|
||||
await update_spend_logs_job(
|
||||
prisma_client=mock_prisma_client,
|
||||
db_writer_client=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert {
|
||||
"rollup_transactions": mock_prisma_client.db.batch_.call_count,
|
||||
"queue_after": mock_prisma_client.tool_usage_transactions,
|
||||
"autorouter_queue_after": mock_prisma_client.autorouter_turn_transactions,
|
||||
} == {"rollup_transactions": 2, "queue_after": [first, second], "autorouter_queue_after": []}
|
||||
|
||||
|
||||
def _model_usage_transaction(model: str) -> object:
|
||||
from litellm.proxy.db.model_usage_rollup import ModelUsageKey, ModelUsageTransaction
|
||||
|
||||
return ModelUsageTransaction(
|
||||
key=ModelUsageKey(
|
||||
date="2026-10-03", model_group="gpt", model=model, custom_llm_provider="openai", task_type="chat"
|
||||
),
|
||||
spend=0.001,
|
||||
prompt_tokens=1,
|
||||
completion_tokens=1,
|
||||
successful=True,
|
||||
)
|
||||
|
||||
|
||||
def _model_usage_batch_failing_with(error: Exception, mock_prisma_client: MagicMock) -> MagicMock:
|
||||
"""Point the model usage rollup's ``batch_()`` at a transaction that fails on commit, with
|
||||
no tool usage queued so the tool rollup never opens one."""
|
||||
mock_prisma_client.db.litellm_spendlogtoolindex.create_many = AsyncMock()
|
||||
mock_prisma_client.tool_usage_transactions = []
|
||||
mock_prisma_client.spend_log_transactions = []
|
||||
rollup_batch = MagicMock()
|
||||
rollup_batch.__aenter__ = AsyncMock(return_value=MagicMock())
|
||||
rollup_batch.__aexit__ = AsyncMock(side_effect=error)
|
||||
mock_prisma_client.db.batch_ = MagicMock(return_value=rollup_batch)
|
||||
return rollup_batch
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_usage_flush_requeues_and_carries_on_when_its_row_lock_times_out(
|
||||
mock_prisma_client: MagicMock,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""The DailyModelUsage rollup shares its rows across every pod exactly like the tool rollup.
|
||||
A batch cancelled under ``lock_timeout`` never applied, so its transactions go back to the
|
||||
queue head for the next flush instead of being dropped, and because only that rollup's row
|
||||
was contended the job still drains the auto-router queue behind it."""
|
||||
monkeypatch.setattr("litellm.proxy.db.model_usage_rollup.asyncio.sleep", AsyncMock(return_value=None))
|
||||
_model_usage_batch_failing_with(_postgres_lock_timeout(), mock_prisma_client)
|
||||
first, second = _model_usage_transaction("gpt-4o"), _model_usage_transaction("gpt-4o-mini")
|
||||
mock_prisma_client.model_usage_transactions = [first, second]
|
||||
mock_prisma_client.autorouter_turn_transactions = [MagicMock(name="autorouter_turn")]
|
||||
|
||||
await update_spend_logs_job(
|
||||
prisma_client=mock_prisma_client,
|
||||
db_writer_client=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert {
|
||||
"queue_after": mock_prisma_client.model_usage_transactions,
|
||||
"autorouter_queue_after": mock_prisma_client.autorouter_turn_transactions,
|
||||
} == {"queue_after": [first, second], "autorouter_queue_after": []}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_usage_flush_requeues_and_stops_the_job_when_postgres_is_out_of_connections(
|
||||
mock_prisma_client: MagicMock,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""A batch refused for want of a connection (53300) was never sent either, so it is kept, but
|
||||
a full server will refuse the writers behind it too: the job stops there like the tool
|
||||
rollup does, so the drain loop backs off instead of re-hitting the server."""
|
||||
monkeypatch.setattr("litellm.proxy.db.model_usage_rollup.asyncio.sleep", AsyncMock(return_value=None))
|
||||
_model_usage_batch_failing_with(_postgres_out_of_connections(), mock_prisma_client)
|
||||
first, second = _model_usage_transaction("gpt-4o"), _model_usage_transaction("gpt-4o-mini")
|
||||
mock_prisma_client.model_usage_transactions = [first, second]
|
||||
mock_prisma_client.autorouter_turn_transactions = [MagicMock(name="autorouter_turn")]
|
||||
|
||||
with pytest.raises(Exception, match="too many clients already"):
|
||||
await update_spend_logs_job(
|
||||
prisma_client=mock_prisma_client,
|
||||
db_writer_client=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert {
|
||||
"queue_after": mock_prisma_client.model_usage_transactions,
|
||||
"autorouter_queue_after": len(mock_prisma_client.autorouter_turn_transactions),
|
||||
} == {"queue_after": [first, second], "autorouter_queue_after": 1}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_usage_flush_still_drops_a_batch_that_may_have_applied(
|
||||
mock_prisma_client: MagicMock,
|
||||
) -> None:
|
||||
"""Any other failure is ambiguous about whether the increments landed, so replaying it could
|
||||
double-count: the batch is logged and dropped, and the job carries on."""
|
||||
_model_usage_batch_failing_with(RuntimeError("engine lost mid-commit"), mock_prisma_client)
|
||||
mock_prisma_client.model_usage_transactions = [_model_usage_transaction("gpt-4o")]
|
||||
|
||||
await update_spend_logs_job(
|
||||
prisma_client=mock_prisma_client,
|
||||
db_writer_client=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert mock_prisma_client.model_usage_transactions == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue