Revert "fix(proxy): durable spend-log flush with redis buffer and passthrough attribution"

This reverts commit 9a3a16fb53.
This commit is contained in:
mubashir1osmani 2026-06-08 18:31:43 -07:00
parent 9a3a16fb53
commit 586e31f1c4
14 changed files with 136 additions and 1310 deletions

View file

@ -295,7 +295,6 @@ REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY = (
)
REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_agent_spend_update_buffer"
REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY = "litellm_daily_tag_spend_update_buffer"
REDIS_SPEND_LOG_BUFFER_KEY = "litellm_spend_log_buffer"
MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", 100))
# Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth
LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000))
@ -1495,15 +1494,6 @@ SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float(
)
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))
SPEND_LOG_QUEUE_POLL_INTERVAL = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0))
DEFAULT_SPEND_LOG_FLUSH_MAX_RETRIES = 3
SPEND_LOG_FLUSH_MAX_RETRIES = max(
0,
int(
os.getenv(
"SPEND_LOG_FLUSH_MAX_RETRIES", DEFAULT_SPEND_LOG_FLUSH_MAX_RETRIES
)
),
)
SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = int(
os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)
)

View file

@ -2585,10 +2585,6 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.",
)
spend_log_flush_max_retries: Optional[int] = Field(
None,
description="Max retries when flushing spend logs to the DB on transient errors (deadlock, pool timeout, connection errors). Defaults to SPEND_LOG_FLUSH_MAX_RETRIES env (3).",
)
mcp_internal_ip_ranges: Optional[List[str]] = Field(
None,
description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).",
@ -3639,8 +3635,6 @@ class SpendLogsMetadata(TypedDict):
cost_breakdown: Optional[
CostBreakdown
] # Detailed cost breakdown (input_cost, output_cost, margin, discount, etc.)
user_api_key_request_route: Optional[str]
passthrough_target_url: Optional[str]
class SpendLogsPayload(TypedDict):
@ -3966,42 +3960,6 @@ DB_CONNECTION_ERROR_TYPES = (
)
def is_spend_log_flush_retryable_error(e: Exception) -> bool:
import asyncio
import prisma.errors
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
return True
if isinstance(
e,
(
httpx.PoolTimeout,
httpx.TimeoutException,
httpx.WriteTimeout,
httpx.ConnectTimeout,
asyncio.TimeoutError,
),
):
return True
if isinstance(e, prisma.errors.PrismaError):
error_message = str(e).lower()
retry_keywords = (
"deadlock",
"could not serialize",
"p2034",
"pool timeout",
"timed out",
"timeout",
"connection",
"too many connections",
"server closed",
"transaction failed",
)
return any(keyword in error_message for keyword in retry_keywords)
return False
class SSOUserDefinedValues(TypedDict):
models: List[str]
user_id: str

View file

@ -55,9 +55,6 @@ from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
)
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
from litellm.proxy.db.db_transaction_queue.spend_log_redis_buffer import (
SpendLogRedisBuffer,
)
from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue
from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import (
ToolDiscoveryQueue,
@ -115,7 +112,6 @@ class DBSpendUpdateWriter:
):
self.redis_cache = redis_cache
self.redis_update_buffer = RedisUpdateBuffer(redis_cache=self.redis_cache)
self.spend_log_redis_buffer = SpendLogRedisBuffer(redis_cache=self.redis_cache)
self.pod_lock_manager = PodLockManager()
self.spend_update_queue = SpendUpdateQueue()
self.tool_discovery_queue = ToolDiscoveryQueue()
@ -759,12 +755,12 @@ class DBSpendUpdateWriter:
payload.get("request_id"), payload.get("spend")
)
)
if prisma_client is not None:
if prisma_client is not None and spend_logs_url is not None:
async with prisma_client._spend_log_transactions_lock:
prisma_client.spend_log_transactions.append(payload)
elif prisma_client is not None:
async with prisma_client._spend_log_transactions_lock:
prisma_client.spend_log_transactions.append(payload)
await self.spend_log_redis_buffer.buffer_spend_log_row(
cast(SpendLogsPayload, payload)
)
else:
verbose_proxy_logger.debug(
"prisma_client is None. Skipping writing spend logs to db."

View file

@ -1,96 +0,0 @@
from typing import TYPE_CHECKING, Dict, List, Optional, Union
from litellm._logging import verbose_proxy_logger
from litellm.constants import REDIS_SPEND_LOG_BUFFER_KEY
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.proxy._types import SpendLogsPayload
from litellm.secret_managers.main import str_to_bool
if TYPE_CHECKING:
from litellm.caching.redis_cache import RedisCache
class SpendLogRedisBuffer:
def __init__(self, redis_cache: Optional["RedisCache"] = None):
self.redis_cache = redis_cache
def is_enabled(self) -> bool:
from typing import Union as TypingUnion
from litellm.proxy.proxy_server import general_settings
if self.redis_cache is None:
return False
_use_redis_transaction_buffer: Optional[TypingUnion[bool, str]] = (
general_settings.get("use_redis_transaction_buffer", False)
)
if isinstance(_use_redis_transaction_buffer, str):
_use_redis_transaction_buffer = str_to_bool(
_use_redis_transaction_buffer
)
if _use_redis_transaction_buffer is None:
return False
return _use_redis_transaction_buffer
async def buffer_spend_log_row(self, payload: SpendLogsPayload) -> None:
if not self.is_enabled() or self.redis_cache is None:
return
try:
await self.redis_cache.async_rpush(
REDIS_SPEND_LOG_BUFFER_KEY,
[safe_dumps(payload)],
)
except Exception as e:
verbose_proxy_logger.exception(
"SpendLogRedisBuffer: failed to buffer spend log row: %s", e
)
async def pop_buffered_spend_log_rows(
self, max_rows: int
) -> List[SpendLogsPayload]:
if not self.is_enabled() or self.redis_cache is None or max_rows <= 0:
return []
rows: List[SpendLogsPayload] = []
try:
for _ in range(max_rows):
serialized_payload = await self.redis_cache.async_lpop(
REDIS_SPEND_LOG_BUFFER_KEY
)
if serialized_payload is None:
break
parsed_payload = safe_json_loads(serialized_payload)
if isinstance(parsed_payload, dict):
rows.append(parsed_payload)
except Exception as e:
verbose_proxy_logger.exception(
"SpendLogRedisBuffer: failed to pop buffered spend log rows: %s", e
)
return rows
async def requeue_spend_log_rows(self, rows: List[SpendLogsPayload]) -> None:
if not rows:
return
if not self.is_enabled() or self.redis_cache is None:
return
try:
for payload in rows:
await self.redis_cache.async_rpush(
REDIS_SPEND_LOG_BUFFER_KEY,
[safe_dumps(payload)],
)
except Exception as e:
verbose_proxy_logger.exception(
"SpendLogRedisBuffer: failed to requeue spend log rows: %s", e
)
async def get_buffered_row_count(self) -> int:
if not self.is_enabled() or self.redis_cache is None:
return 0
try:
return await self.redis_cache.async_llen(REDIS_SPEND_LOG_BUFFER_KEY)
except Exception:
return 0

View file

@ -557,8 +557,6 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
request=request,
metadata=_metadata,
)
_metadata["passthrough_target_url"] = passthrough_logging_payload.get("url")
_metadata["user_api_key_request_route"] = user_api_key_dict.request_route
# Set internal keys after merging client-supplied metadata so a request
# body that mirrors them cannot clobber the authenticated key or the
@ -1304,7 +1302,6 @@ async def pass_through_request( # noqa: PLR0915
model_id=None,
cache_key=None,
api_base=str(url._uri_reference),
response_cost=cost_per_request,
)
response_headers = HttpPassThroughEndpointHelpers.get_response_headers(

View file

@ -42,7 +42,6 @@ from litellm.proxy._types import (
ProxyException,
SpendLogsMetadata,
SpendLogsPayload,
is_spend_log_flush_retryable_error,
)
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.types.guardrails import GuardrailEventHooks
@ -5341,7 +5340,6 @@ class ProxyUpdateSpend:
"Spend tracking - processing %d spend logs for DB write",
len(logs_to_process),
)
spend_log_redis_buffer = _get_spend_log_redis_buffer(proxy_logging_obj)
start_time = time.time()
try:
for i in range(n_retry_times + 1):
@ -5363,6 +5361,7 @@ class ProxyUpdateSpend:
)
del json_data
if response.status_code == 200:
# Items already removed from queue at start of function
pass
else:
for j in range(0, len(logs_to_process), BATCH_SIZE):
@ -5377,37 +5376,33 @@ class ProxyUpdateSpend:
verbose_proxy_logger.debug(
f"Flushed {len(batch)} logs to the DB."
)
# Explicitly clear batch memory
del batch, batch_with_dates
# Items already removed from queue at start of function
async with prisma_client._spend_log_transactions_lock:
remaining_count = len(prisma_client.spend_log_transactions)
verbose_proxy_logger.debug(
f"{len(logs_to_process)} logs processed. Remaining in queue: {remaining_count}"
)
break
except Exception as e:
if (
is_spend_log_flush_retryable_error(e)
and i < n_retry_times
):
verbose_proxy_logger.warning(
"Spend tracking - retryable error 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(2**i)
continue
raise
except DB_CONNECTION_ERROR_TYPES as e:
if i is None:
i = 0
verbose_proxy_logger.warning(
"Spend tracking - DB connection error writing spend logs, "
"retry %d/%d. logs_count=%d, error=%s",
i + 1,
n_retry_times,
len(logs_to_process),
str(e),
)
if i >= n_retry_times:
raise
await asyncio.sleep(2**i)
except Exception as e:
if is_spend_log_flush_retryable_error(e):
await _requeue_failed_spend_logs(
prisma_client=prisma_client,
logs_to_process=logs_to_process,
spend_log_redis_buffer=spend_log_redis_buffer,
)
# Logs already removed from queue at start - don't put them back
# This matches the original behavior where logs are removed even on error
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
)
@ -5521,107 +5516,6 @@ async def update_daily_tag_spend(
verbose_proxy_logger.error(f"Error updating daily tag spend: {e}")
def get_spend_log_flush_max_retries() -> int:
"""
Max retries when flushing spend logs to the DB on transient errors.
Precedence (highest wins):
1. general_settings.spend_log_flush_max_retries
2. SPEND_LOG_FLUSH_MAX_RETRIES env (via litellm.constants)
3. DEFAULT_SPEND_LOG_FLUSH_MAX_RETRIES (3)
"""
from litellm.constants import (
DEFAULT_SPEND_LOG_FLUSH_MAX_RETRIES,
SPEND_LOG_FLUSH_MAX_RETRIES,
)
default_retries = SPEND_LOG_FLUSH_MAX_RETRIES or DEFAULT_SPEND_LOG_FLUSH_MAX_RETRIES
try:
from litellm.proxy.proxy_server import general_settings
except Exception:
return default_retries
configured = general_settings.get("spend_log_flush_max_retries")
if configured is None:
return default_retries
return max(0, int(configured))
def _dedupe_spend_logs_by_request_id(
logs: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
deduped: List[Dict[str, Any]] = []
seen_request_ids: set = set()
for log in logs:
request_id = log.get("request_id")
if request_id is not None and request_id in seen_request_ids:
continue
if request_id is not None:
seen_request_ids.add(request_id)
deduped.append(log)
return deduped
def _get_spend_log_redis_buffer(
proxy_logging_obj: ProxyLogging,
) -> Optional[Any]:
db_spend_update_writer = getattr(
proxy_logging_obj, "db_spend_update_writer", None
)
if db_spend_update_writer is None:
return None
return getattr(db_spend_update_writer, "spend_log_redis_buffer", None)
async def _requeue_failed_spend_logs(
prisma_client: PrismaClient,
logs_to_process: List[Dict[str, Any]],
spend_log_redis_buffer: Optional[Any],
) -> None:
if not logs_to_process:
return
async with prisma_client._spend_log_transactions_lock:
prisma_client.spend_log_transactions = (
logs_to_process + prisma_client.spend_log_transactions
)
if spend_log_redis_buffer is not None:
if spend_log_redis_buffer.is_enabled() is True:
await spend_log_redis_buffer.requeue_spend_log_rows(
[cast(SpendLogsPayload, log) for log in logs_to_process]
)
verbose_proxy_logger.warning(
"Spend tracking - re-queued %d spend logs after flush failure",
len(logs_to_process),
)
async def _collect_spend_logs_for_flush(
prisma_client: PrismaClient,
proxy_logging_obj: ProxyLogging,
max_logs: int,
) -> List[Dict[str, Any]]:
spend_log_redis_buffer = _get_spend_log_redis_buffer(proxy_logging_obj)
redis_logs: List[Dict[str, Any]] = []
remaining_slots = max_logs
if spend_log_redis_buffer is not None:
if spend_log_redis_buffer.is_enabled() is True:
redis_logs = await spend_log_redis_buffer.pop_buffered_spend_log_rows(
max_rows=remaining_slots
)
remaining_slots = max(0, max_logs - len(redis_logs))
memory_logs: List[Dict[str, Any]] = []
async with prisma_client._spend_log_transactions_lock:
if remaining_slots > 0:
memory_logs = prisma_client.spend_log_transactions[:remaining_slots]
prisma_client.spend_log_transactions = (
prisma_client.spend_log_transactions[len(memory_logs) :]
)
return _dedupe_spend_logs_by_request_id(redis_logs + memory_logs)
async def update_spend_logs_job(
prisma_client: PrismaClient,
db_writer_client: Optional[AsyncHTTPHandler],
@ -5633,29 +5527,20 @@ async def update_spend_logs_job(
This job is triggered based on queue size rather than time.
Pops the batch once, writes spend logs, then runs guardrail usage tracking.
"""
# Retries on transient DB errors (deadlock, pool timeout, etc.).
# Defaults to 3; override via general_settings.spend_log_flush_max_retries
# or SPEND_LOG_FLUSH_MAX_RETRIES env.
n_retry_times = get_spend_log_flush_max_retries()
n_retry_times = 3
MAX_LOGS_PER_INTERVAL = 10000
spend_log_redis_buffer = _get_spend_log_redis_buffer(proxy_logging_obj)
redis_queue_size = 0
if spend_log_redis_buffer is not None and spend_log_redis_buffer.is_enabled() is True:
redis_queue_size = await spend_log_redis_buffer.get_buffered_row_count()
# Atomically pop batch from queue
async with prisma_client._spend_log_transactions_lock:
queue_size = len(prisma_client.spend_log_transactions)
if queue_size == 0:
return
async with prisma_client._spend_log_transactions_lock:
memory_queue_size = len(prisma_client.spend_log_transactions)
if memory_queue_size == 0 and redis_queue_size == 0:
return
logs_to_process = await _collect_spend_logs_for_flush(
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
max_logs=MAX_LOGS_PER_INTERVAL,
)
if not logs_to_process:
return
logs_to_process = prisma_client.spend_log_transactions[:MAX_LOGS_PER_INTERVAL]
prisma_client.spend_log_transactions = prisma_client.spend_log_transactions[
len(logs_to_process) :
]
await ProxyUpdateSpend.update_spend_logs(
n_retry_times=n_retry_times,

View file

@ -1,26 +0,0 @@
from __future__ import annotations
import asyncio
import httpx
import prisma.errors
import pytest
from litellm.proxy._types import is_spend_log_flush_retryable_error
@pytest.mark.parametrize(
"error,expected",
[
(httpx.ConnectError("connect"), True),
(httpx.ReadError("read"), True),
(httpx.ReadTimeout("read timeout"), True),
(httpx.PoolTimeout("pool timeout"), True),
(asyncio.TimeoutError(), True),
(prisma.errors.PrismaError("deadlock detected"), True),
(prisma.errors.PrismaError("pool timeout waiting for connection"), True),
(ValueError("bad data"), False),
],
)
def test_is_spend_log_flush_retryable_error(error: Exception, expected: bool) -> None:
assert is_spend_log_flush_retryable_error(error) is expected

View file

@ -1,128 +0,0 @@
from __future__ import annotations
import json
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.constants import REDIS_SPEND_LOG_BUFFER_KEY
from litellm.proxy.db.db_transaction_queue.spend_log_redis_buffer import (
SpendLogRedisBuffer,
)
@pytest.fixture
def mock_redis_cache() -> MagicMock:
cache = MagicMock()
cache.redis_batch_writing_buffer_key = "test-buffer-key"
cache.async_rpush = AsyncMock()
cache.async_lpop = AsyncMock(return_value=None)
cache.async_llen = AsyncMock(return_value=0)
return cache
@pytest.mark.asyncio
async def test_buffer_spend_log_row_pushes_to_redis_when_enabled(
mock_redis_cache: MagicMock,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"use_redis_transaction_buffer": True},
)
buffer = SpendLogRedisBuffer(redis_cache=mock_redis_cache)
payload = {"request_id": "req-1", "spend": 12.0}
await buffer.buffer_spend_log_row(payload) # type: ignore[arg-type]
mock_redis_cache.async_rpush.assert_awaited_once()
assert mock_redis_cache.async_rpush.await_args.args[0] == REDIS_SPEND_LOG_BUFFER_KEY
@pytest.mark.asyncio
async def test_pop_buffered_spend_log_rows_returns_deserialized_rows(
mock_redis_cache: MagicMock,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"use_redis_transaction_buffer": True},
)
payload = {"request_id": "req-2", "spend": 3.5}
serialized = safe_dumps(payload)
mock_redis_cache.async_lpop = AsyncMock(side_effect=[serialized, None])
buffer = SpendLogRedisBuffer(redis_cache=mock_redis_cache)
rows = await buffer.pop_buffered_spend_log_rows(max_rows=5)
assert rows == [payload]
@pytest.mark.asyncio
async def test_requeue_spend_log_rows_pushes_each_row(
mock_redis_cache: MagicMock,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"use_redis_transaction_buffer": True},
)
buffer = SpendLogRedisBuffer(redis_cache=mock_redis_cache)
rows = [
{"request_id": "req-a", "spend": 1.0},
{"request_id": "req-b", "spend": 2.0},
]
await buffer.requeue_spend_log_rows(rows) # type: ignore[arg-type]
assert mock_redis_cache.async_rpush.await_count == 2
@pytest.mark.asyncio
async def test_buffer_spend_log_row_noop_when_disabled(
mock_redis_cache: MagicMock,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"use_redis_transaction_buffer": False},
)
buffer = SpendLogRedisBuffer(redis_cache=mock_redis_cache)
await buffer.buffer_spend_log_row({"request_id": "req-x", "spend": 1.0}) # type: ignore[arg-type]
mock_redis_cache.async_rpush.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_buffered_row_count_returns_redis_length(
mock_redis_cache: MagicMock,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"use_redis_transaction_buffer": True},
)
mock_redis_cache.async_llen = AsyncMock(return_value=7)
buffer = SpendLogRedisBuffer(redis_cache=mock_redis_cache)
assert await buffer.get_buffered_row_count() == 7

View file

@ -1656,68 +1656,3 @@ async def test_commit_spend_updates_iterates_in_sorted_order(
)
assert captured_where_values == expected_order
@pytest.mark.asyncio
async def test_insert_spend_log_to_db_buffers_memory_and_redis(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"use_redis_transaction_buffer": True},
)
mock_redis_cache = MagicMock()
mock_redis_cache.async_rpush = AsyncMock()
mock_prisma_client = MagicMock()
mock_prisma_client.spend_log_transactions = []
mock_prisma_client._spend_log_transactions_lock = asyncio.Lock()
db_writer = DBSpendUpdateWriter(redis_cache=mock_redis_cache)
payload = {
"request_id": "req-buffer-test",
"spend": 0.42,
"model": "gpt-4o-mini",
}
await db_writer._insert_spend_log_to_db(
payload=payload,
prisma_client=mock_prisma_client,
)
assert mock_prisma_client.spend_log_transactions == [payload]
mock_redis_cache.async_rpush.assert_awaited_once()
@pytest.mark.asyncio
async def test_insert_spend_log_to_db_skips_redis_when_buffer_disabled(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"use_redis_transaction_buffer": False},
)
mock_redis_cache = MagicMock()
mock_redis_cache.async_rpush = AsyncMock()
mock_prisma_client = MagicMock()
mock_prisma_client.spend_log_transactions = []
mock_prisma_client._spend_log_transactions_lock = asyncio.Lock()
db_writer = DBSpendUpdateWriter(redis_cache=mock_redis_cache)
payload = {"request_id": "req-no-redis", "spend": 0.1}
await db_writer._insert_spend_log_to_db(
payload=payload,
prisma_client=mock_prisma_client,
)
assert mock_prisma_client.spend_log_transactions == [payload]
mock_redis_cache.async_rpush.assert_not_awaited()

View file

@ -1,350 +0,0 @@
import asyncio
import json
import os
import sys
from contextlib import ExitStack
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import httpx
import litellm
import pytest
from starlette.requests import Request
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
pass_through_request,
)
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
from litellm.proxy.utils import hash_token
_PT_MOD = "litellm.proxy.pass_through_endpoints.pass_through_endpoints"
async def _mock_upstream_request(*args, **kwargs):
mock_response = httpx.Response(200, json={})
mock_response.request = Mock(spec=httpx.Request)
return mock_response
def _bria_request() -> Request:
return Request(
{
"type": "http",
"method": "POST",
"path": "/bria",
"raw_path": b"/bria",
"query_string": b"",
"headers": [
(b"content-type", b"application/json"),
(b"x-api-key", b"dummy-api-key"),
],
"scheme": "http",
"server": ("testserver", 80),
"client": ("testclient", 50000),
}
)
def _threads_request() -> Request:
return Request(
{
"type": "http",
"method": "POST",
"path": "/v1/threads",
"raw_path": b"/v1/threads",
"query_string": b"",
"headers": [
(b"content-type", b"application/json"),
(b"authorization", b"Bearer sk-test"),
(b"openai-beta", b"assistants=v2"),
],
"scheme": "http",
"server": ("testserver", 80),
"client": ("testclient", 50000),
}
)
def _header_test_patches(mock_proxy_logging):
mock_async_client = AsyncMock()
mock_async_client_obj = MagicMock()
mock_async_client_obj.client = mock_async_client
mock_async_client.request = AsyncMock(side_effect=_mock_upstream_request)
mock_pt_logging = MagicMock()
mock_pt_logging.pass_through_async_success_handler = AsyncMock()
patches = [
patch(
f"{_PT_MOD}.HttpPassThroughEndpointHelpers.non_streaming_http_request_handler",
new_callable=AsyncMock,
side_effect=_mock_upstream_request,
),
patch(f"{_PT_MOD}._is_streaming_response", return_value=False),
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging),
patch(f"{_PT_MOD}.pass_through_endpoint_logging", mock_pt_logging),
patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj),
patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}),
patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}),
]
stack = ExitStack()
for p in patches:
stack.enter_context(p)
return stack
def _e2e_spend_tracking_patches(
mock_proxy_logging,
captured_db_calls: list,
captured_spend_payloads: list,
captured_counter_calls: list,
background_tasks: list,
):
mock_async_client = AsyncMock()
mock_async_client_obj = MagicMock()
mock_async_client_obj.client = mock_async_client
mock_async_client.request = AsyncMock(side_effect=_mock_upstream_request)
mock_db_writer = MagicMock()
async def capture_update_database(**kwargs):
payload = get_logging_payload(
kwargs=kwargs["kwargs"],
response_obj=kwargs["completion_response"],
start_time=kwargs["start_time"],
end_time=kwargs["end_time"],
)
payload["spend"] = kwargs["response_cost"] or 0.0
captured_db_calls.append(kwargs)
captured_spend_payloads.append(payload)
mock_db_writer.update_database = AsyncMock(side_effect=capture_update_database)
mock_proxy_logging.db_spend_update_writer = mock_db_writer
mock_proxy_logging.failed_tracking_alert = AsyncMock()
mock_proxy_logging.slack_alerting_instance = MagicMock()
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
async def capture_increment_spend_counters(**kwargs):
captured_counter_calls.append(kwargs)
real_create_task = asyncio.create_task
def capture_create_task(coro):
task = real_create_task(coro)
background_tasks.append(task)
return task
patches = [
patch(
f"{_PT_MOD}.HttpPassThroughEndpointHelpers.non_streaming_http_request_handler",
new_callable=AsyncMock,
side_effect=_mock_upstream_request,
),
patch(f"{_PT_MOD}._is_streaming_response", return_value=False),
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging),
patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj),
patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}),
patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}),
patch(
"litellm.proxy.proxy_server.increment_spend_counters",
side_effect=capture_increment_spend_counters,
),
patch("litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock),
patch("asyncio.create_task", side_effect=capture_create_task),
]
stack = ExitStack()
for p in patches:
stack.enter_context(p)
return stack
@pytest.mark.asyncio
async def test_bria_passthrough_returns_configured_response_cost_header():
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_success_hook = AsyncMock(return_value={})
with _header_test_patches(mock_proxy_logging):
response = await pass_through_request(
request=_bria_request(),
target="https://engine.prod.bria-api.com",
custom_headers={"x-api-key": "dummy-api-key"},
user_api_key_dict=UserAPIKeyAuth(
api_key="hashed-test-key",
user_id="test-user",
team_id="test-team",
),
cost_per_request=12.0,
)
assert response.status_code == 200
assert response.body == b"{}"
assert "x-litellm-response-cost" in response.headers
assert float(response.headers["x-litellm-response-cost"]) == 12.0
@pytest.mark.asyncio
async def test_bria_passthrough_cost_per_request_e2e_spend_tracking():
hashed_api_key = hash_token("sk-bria-spend-test")
user_api_key_dict = UserAPIKeyAuth(
api_key=hashed_api_key,
user_id="test-user",
team_id="test-team",
request_route="/bria",
)
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_success_hook = AsyncMock(return_value={})
captured_db_calls: list = []
captured_spend_payloads: list = []
captured_counter_calls: list = []
background_tasks: list = []
original_callbacks = list(litellm._async_success_callback)
litellm._async_success_callback = [_ProxyDBLogger()]
try:
with _e2e_spend_tracking_patches(
mock_proxy_logging,
captured_db_calls,
captured_spend_payloads,
captured_counter_calls,
background_tasks,
):
response = await pass_through_request(
request=_bria_request(),
target="https://engine.prod.bria-api.com",
custom_headers={"x-api-key": "dummy-api-key"},
user_api_key_dict=user_api_key_dict,
cost_per_request=12.0,
)
await asyncio.gather(*background_tasks, return_exceptions=True)
mock_proxy_logging.failed_tracking_alert.assert_not_awaited()
finally:
litellm._async_success_callback = original_callbacks
assert response.status_code == 200
assert float(response.headers["x-litellm-response-cost"]) == 12.0
assert len(captured_db_calls) == 1
db_call = captured_db_calls[0]
assert db_call["response_cost"] == 12.0
assert db_call["token"] == hashed_api_key
assert db_call["user_id"] == "test-user"
assert db_call["team_id"] == "test-team"
standard_logging_object = db_call["kwargs"]["standard_logging_object"]
assert standard_logging_object is not None
assert standard_logging_object["response_cost"] == 12.0
assert standard_logging_object["call_type"] == "pass_through_endpoint"
assert len(captured_spend_payloads) == 1
spend_payload = captured_spend_payloads[0]
assert spend_payload["spend"] == 12.0
assert spend_payload["api_key"] == hashed_api_key
spend_metadata = spend_payload["metadata"]
if isinstance(spend_metadata, str):
spend_metadata = json.loads(spend_metadata)
assert spend_metadata["user_api_key_request_route"] == "/bria"
assert (
spend_metadata["passthrough_target_url"]
== "https://engine.prod.bria-api.com"
)
assert len(captured_counter_calls) == 1
counter_call = captured_counter_calls[0]
assert counter_call["response_cost"] == 12.0
assert counter_call["token"] == hashed_api_key
assert counter_call["user_id"] == "test-user"
assert counter_call["team_id"] == "test-team"
@pytest.mark.asyncio
async def test_threads_passthrough_cost_per_request_e2e_spend_tracking():
hashed_api_key = hash_token("sk-threads-spend-test")
user_api_key_dict = UserAPIKeyAuth(
api_key=hashed_api_key,
user_id="threads-user",
team_id="threads-team",
request_route="/v1/threads",
)
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_success_hook = AsyncMock(return_value={})
captured_db_calls: list = []
captured_spend_payloads: list = []
captured_counter_calls: list = []
background_tasks: list = []
original_callbacks = list(litellm._async_success_callback)
litellm._async_success_callback = [_ProxyDBLogger()]
try:
with _e2e_spend_tracking_patches(
mock_proxy_logging,
captured_db_calls,
captured_spend_payloads,
captured_counter_calls,
background_tasks,
):
response = await pass_through_request(
request=_threads_request(),
target="https://api.openai.com/v1/threads",
custom_headers={
"Content-Type": "application/json",
"Authorization": "Bearer sk-test",
"OpenAI-Beta": "assistants=v2",
},
user_api_key_dict=user_api_key_dict,
cost_per_request=0.05,
)
await asyncio.gather(*background_tasks, return_exceptions=True)
mock_proxy_logging.failed_tracking_alert.assert_not_awaited()
finally:
litellm._async_success_callback = original_callbacks
assert response.status_code == 200
assert float(response.headers["x-litellm-response-cost"]) == 0.05
assert len(captured_db_calls) == 1
db_call = captured_db_calls[0]
assert db_call["response_cost"] == 0.05
assert db_call["token"] == hashed_api_key
assert db_call["user_id"] == "threads-user"
assert db_call["team_id"] == "threads-team"
standard_logging_object = db_call["kwargs"]["standard_logging_object"]
assert standard_logging_object is not None
assert standard_logging_object["response_cost"] == 0.05
assert standard_logging_object["call_type"] == "pass_through_endpoint"
assert len(captured_spend_payloads) == 1
spend_payload = captured_spend_payloads[0]
assert spend_payload["spend"] == 0.05
assert spend_payload["api_key"] == hashed_api_key
spend_metadata = spend_payload["metadata"]
if isinstance(spend_metadata, str):
spend_metadata = json.loads(spend_metadata)
assert spend_metadata["user_api_key_request_route"] == "/v1/threads"
assert spend_metadata["passthrough_target_url"] == "https://api.openai.com/v1/threads"
assert len(captured_counter_calls) == 1
counter_call = captured_counter_calls[0]
assert counter_call["response_cost"] == 0.05
assert counter_call["token"] == hashed_api_key
assert counter_call["user_id"] == "threads-user"
assert counter_call["team_id"] == "threads-team"

View file

@ -0,0 +1,102 @@
# Spend Tracking Test Runbook
This runbook documents the regression coverage that guards litellm's spend/cost
tracking against silent breakage before a release. It has two halves: a set of
deterministic offline unit tests (Part A) that run in CI, and an operator-driven
live end-to-end check (Part B) that hits real provider APIs and asserts real
SpendLogs rows. Spend tracking is already heavily tested across the codebase
(cost calculators, SpendLogsPayload construction, the proxy cost callback, the DB
spend-update writer, daily-spend queues); the tests below fill the specific
high-value gaps that previously had no direct coverage.
## Part A: offline regression tests (CI)
Run them with:
```bash
uv run --no-sync python -m pytest \
tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py \
tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py \
tests/test_litellm/test_cost_calculator.py::test_embedding_completion_cost_uses_input_cost_per_token \
-q
```
Each test is written to fail if the specific guard it covers is mutated; the
"What it guards" column names the source line that, when broken, turns the test
red. Verified by hand-mutation (revert each guard and watch the row go red).
| Test name | What it guards | Status |
| --- | --- | --- |
| `TestGetStatusForSpendLog::test_missing_status_key_defaults_to_success` | `_get_status_for_spend_log` default branch (spend_tracking_utils.py) returns success when no status set | pass |
| `TestGetStatusForSpendLog::test_explicit_success_returns_success` | same helper, explicit success preserved | pass |
| `TestGetStatusForSpendLog::test_failure_returns_failure` | the `== "failure"` guard so failed requests are logged as failures | pass |
| `TestGetStatusForSpendLog::test_non_failure_value_returns_success` | kills the "any non-None status -> failure" mutant | pass |
| `test_get_logging_payload_cache_hit_appends_unique_suffix_to_request_id` | cache-hit `_cache_hit{time}` suffix on request_id; without it SpendLogs hits duplicate-key collisions | pass |
| `test_get_logging_payload_failure_status_and_zero_spend` | status wiring at the `get_logging_payload` call site plus `spend` sourced from `response_cost` | pass |
| `test_get_logging_payload_default_status_success` | default status path through `get_logging_payload` | pass |
| `test_track_cost_callback_zeroes_response_cost_on_cache_hit` | cache-hit cost zeroing in `_PROXY_track_cost_callback`; the anti double-charge guard | pass |
| `test_embedding_completion_cost_uses_input_cost_per_token` | embedding cost = `prompt_tokens * input_cost_per_token`; previously no dedicated embedding cost test | pass |
All 9 pass on `claude/spend-tracking-tests-4gbezl`.
## Part B: live end-to-end check (operator-run, real spend logs)
This proves spend tracking against real provider responses and a real database,
which is the closest mirror of what a customer sees. It needs a Postgres
`DATABASE_URL`, an `OPENAI_API_KEY` in `.env`, and outbound access to
api.openai.com. It uses `gpt-5.4-nano` (current cheap small model as of 2026-06)
and `text-embedding-3-small`, with local caching enabled so a repeated identical
chat request produces a cache hit.
Config: `tests/test_litellm/proxy/spend_tracking/e2e_spend_config.yaml`
1. Point at a database and start the proxy (spend logs require a DB):
```bash
export DATABASE_URL='postgresql://user:pass@localhost:5432/litellm'
python litellm/proxy/proxy_cli.py \
--config tests/test_litellm/proxy/spend_tracking/e2e_spend_config.yaml \
--detailed_debug 2>&1 | tee litellm.log
```
2. First chat call (real cost expected):
```bash
curl -s http://localhost:4000/v1/chat/completions \
-H 'Authorization: Bearer sk-1234' -H 'Content-Type: application/json' \
-d '{"model":"gpt-5.4-nano","messages":[{"role":"user","content":"say hello in one word"}]}' | jq '{id, usage}'
```
3. Identical chat call again to trigger a cache hit (cost should be recorded as 0):
```bash
curl -s http://localhost:4000/v1/chat/completions \
-H 'Authorization: Bearer sk-1234' -H 'Content-Type: application/json' \
-d '{"model":"gpt-5.4-nano","messages":[{"role":"user","content":"say hello in one word"}]}' | jq '{id, usage}'
```
4. Embedding call (real cost expected):
```bash
curl -s http://localhost:4000/v1/embeddings \
-H 'Authorization: Bearer sk-1234' -H 'Content-Type: application/json' \
-d '{"model":"text-embedding-3-small","input":"hello world"}' | jq '{model, usage}'
```
5. Wait for the spend-log flush (the proxy batches writes roughly once a minute),
then read the logs back and assert the three rows:
```bash
sleep 65
curl -s 'http://localhost:4000/spend/logs' -H 'Authorization: Bearer sk-1234' \
| jq '[.[] | {request_id, call_type, model, spend, cache_hit}]'
```
Expected: the chat row has `spend > 0`, the embedding row has `spend > 0`, and the
cache-hit row has `spend == 0` with a `request_id` containing `_cache_hit`.
Status: not run in this sandbox. The container has no `.env`/`OPENAI_API_KEY`, no
`DATABASE_URL`, and egress to api.openai.com is blocked by the environment's
network policy (a billed call is not possible here). Run the steps above on a host
with credentials and network access to fill this status; the offline suite in
Part A is what gates CI

View file

@ -216,7 +216,6 @@ async def test_update_spend_logs_failure_raises_after_retries(
monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep)
failed_logs = [make_spend_log_row(request_id="r1")]
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(
side_effect=httpx.ReadError("network blip")
)
@ -228,78 +227,8 @@ async def test_update_spend_logs_failure_raises_after_retries(
prisma_client=mock_prisma_client,
db_writer_client=None,
proxy_logging_obj=proxy_logging,
logs_to_process=failed_logs,
logs_to_process=[make_spend_log_row(request_id="r1")],
)
assert mock_prisma_client.spend_log_transactions == failed_logs
@pytest.mark.asyncio
async def test_update_spend_logs_retries_on_deadlock_error(
mock_prisma_client: Any,
make_spend_log_row: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.utils as utils_mod
import prisma.errors
async def _fake_sleep(_: float) -> None:
return None
monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep)
create_many = AsyncMock(
side_effect=[
prisma.errors.PrismaError("deadlock detected"),
None,
]
)
mock_prisma_client.db.litellm_spendlogs.create_many = create_many
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
logs = [make_spend_log_row(request_id="r1")]
await ProxyUpdateSpend.update_spend_logs(
n_retry_times=1,
prisma_client=mock_prisma_client,
db_writer_client=None,
proxy_logging_obj=proxy_logging,
logs_to_process=logs,
)
assert create_many.await_count == 2
assert mock_prisma_client.spend_log_transactions == []
@pytest.mark.asyncio
async def test_update_spend_logs_retries_on_pool_timeout(
mock_prisma_client: Any,
make_spend_log_row: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import httpx
import litellm.proxy.utils as utils_mod
async def _fake_sleep(_: float) -> None:
return None
monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep)
create_many = AsyncMock(
side_effect=[
httpx.PoolTimeout("pool timeout"),
None,
]
)
mock_prisma_client.db.litellm_spendlogs.create_many = create_many
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
logs = [make_spend_log_row(request_id="r1")]
await ProxyUpdateSpend.update_spend_logs(
n_retry_times=1,
prisma_client=mock_prisma_client,
db_writer_client=None,
proxy_logging_obj=proxy_logging,
logs_to_process=logs,
)
assert create_many.await_count == 2
def test_disable_spend_updates_reflects_general_settings(

View file

@ -1,263 +0,0 @@
"""Regression tests for durable spend-log flushing and re-queue behavior."""
from __future__ import annotations
from typing import Any, Dict, List
from unittest.mock import AsyncMock, MagicMock
import httpx
import prisma.errors
import pytest
from litellm.proxy.utils import (
ProxyUpdateSpend,
_collect_spend_logs_for_flush,
_dedupe_spend_logs_by_request_id,
_requeue_failed_spend_logs,
update_spend_logs_job,
)
def test_dedupe_spend_logs_by_request_id_prefers_first_occurrence() -> None:
logs = [
{"request_id": "a", "spend": 1.0},
{"request_id": "b", "spend": 2.0},
{"request_id": "a", "spend": 3.0},
{"spend": 4.0},
]
assert _dedupe_spend_logs_by_request_id(logs) == [
{"request_id": "a", "spend": 1.0},
{"request_id": "b", "spend": 2.0},
{"spend": 4.0},
]
@pytest.mark.asyncio
async def test_collect_spend_logs_for_flush_merges_redis_and_memory(
mock_prisma_client: Any,
make_spend_log_row: Any,
) -> None:
mock_prisma_client.spend_log_transactions = [
make_spend_log_row(request_id="mem-1"),
make_spend_log_row(request_id="mem-2"),
]
redis_buffer = MagicMock()
redis_buffer.is_enabled = MagicMock(return_value=True)
redis_buffer.pop_buffered_spend_log_rows = AsyncMock(
return_value=[make_spend_log_row(request_id="redis-1")]
)
proxy_logging = MagicMock()
proxy_logging.db_spend_update_writer = MagicMock()
proxy_logging.db_spend_update_writer.spend_log_redis_buffer = redis_buffer
collected = await _collect_spend_logs_for_flush(
prisma_client=mock_prisma_client,
proxy_logging_obj=proxy_logging,
max_logs=10,
)
assert [row["request_id"] for row in collected] == ["redis-1", "mem-1", "mem-2"]
assert mock_prisma_client.spend_log_transactions == []
redis_buffer.pop_buffered_spend_log_rows.assert_awaited_once_with(max_rows=10)
@pytest.mark.asyncio
async def test_collect_spend_logs_for_flush_dedupes_across_sources(
mock_prisma_client: Any,
make_spend_log_row: Any,
) -> None:
duplicate = make_spend_log_row(request_id="dup", spend=1.0)
mock_prisma_client.spend_log_transactions = [
make_spend_log_row(request_id="mem-only"),
make_spend_log_row(request_id="dup", spend=9.0),
]
redis_buffer = MagicMock()
redis_buffer.is_enabled = MagicMock(return_value=True)
redis_buffer.pop_buffered_spend_log_rows = AsyncMock(return_value=[duplicate])
proxy_logging = MagicMock()
proxy_logging.db_spend_update_writer = MagicMock()
proxy_logging.db_spend_update_writer.spend_log_redis_buffer = redis_buffer
collected = await _collect_spend_logs_for_flush(
prisma_client=mock_prisma_client,
proxy_logging_obj=proxy_logging,
max_logs=10,
)
assert [row["request_id"] for row in collected] == ["dup", "mem-only"]
assert collected[0]["spend"] == 1.0
@pytest.mark.asyncio
async def test_requeue_failed_spend_logs_restores_memory_and_redis(
mock_prisma_client: Any,
make_spend_log_row: Any,
) -> None:
failed_rows = [make_spend_log_row(request_id="failed-1")]
mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="queued")]
redis_buffer = MagicMock()
redis_buffer.is_enabled = MagicMock(return_value=True)
redis_buffer.requeue_spend_log_rows = AsyncMock()
await _requeue_failed_spend_logs(
prisma_client=mock_prisma_client,
logs_to_process=failed_rows,
spend_log_redis_buffer=redis_buffer,
)
assert [row["request_id"] for row in mock_prisma_client.spend_log_transactions] == [
"failed-1",
"queued",
]
redis_buffer.requeue_spend_log_rows.assert_awaited_once_with(failed_rows)
@pytest.mark.asyncio
async def test_update_spend_logs_job_processes_redis_only_queue(
mock_prisma_client: Any,
make_spend_log_row: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.guardrails.usage_tracking as guard_mod
import litellm.proxy.db.spend_log_tool_index as tool_mod
import litellm.proxy.utils as utils_mod
mock_prisma_client.spend_log_transactions = []
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock()
redis_row = make_spend_log_row(request_id="redis-only")
redis_buffer = MagicMock()
redis_buffer.is_enabled = MagicMock(return_value=True)
redis_buffer.get_buffered_row_count = AsyncMock(return_value=1)
redis_buffer.pop_buffered_spend_log_rows = AsyncMock(return_value=[redis_row])
redis_buffer.requeue_spend_log_rows = AsyncMock()
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
proxy_logging.db_spend_update_writer = MagicMock()
proxy_logging.db_spend_update_writer.spend_log_redis_buffer = redis_buffer
monkeypatch.setattr(
guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False
)
monkeypatch.setattr(
tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False
)
await update_spend_logs_job(
prisma_client=mock_prisma_client,
db_writer_client=None,
proxy_logging_obj=proxy_logging,
)
assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == 1
assert (
mock_prisma_client.db.litellm_spendlogs.create_many.await_args.kwargs["data"][0][
"request_id"
]
== "redis-only"
)
redis_buffer.requeue_spend_log_rows.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_spend_logs_job_requeues_after_retryable_flush_failure(
mock_prisma_client: Any,
make_spend_log_row: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.utils as utils_mod
async def _fake_sleep(_: float) -> None:
return None
monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep)
failed_row = make_spend_log_row(request_id="lost-if-not-requeued")
mock_prisma_client.spend_log_transactions = [failed_row]
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(
side_effect=httpx.ReadError("connection reset")
)
redis_buffer = MagicMock()
redis_buffer.is_enabled = MagicMock(return_value=True)
redis_buffer.get_buffered_row_count = AsyncMock(return_value=0)
redis_buffer.pop_buffered_spend_log_rows = AsyncMock(return_value=[])
redis_buffer.requeue_spend_log_rows = AsyncMock()
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
proxy_logging.db_spend_update_writer = MagicMock()
proxy_logging.db_spend_update_writer.spend_log_redis_buffer = redis_buffer
with pytest.raises(httpx.ReadError):
await update_spend_logs_job(
prisma_client=mock_prisma_client,
db_writer_client=None,
proxy_logging_obj=proxy_logging,
)
assert mock_prisma_client.spend_log_transactions == [failed_row]
redis_buffer.requeue_spend_log_rows.assert_awaited_once()
@pytest.mark.asyncio
async def test_update_spend_logs_does_not_requeue_non_retryable_errors(
mock_prisma_client: Any,
make_spend_log_row: Any,
) -> None:
failed_logs = [make_spend_log_row(request_id="bad-row")]
mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(
side_effect=ValueError("invalid spend log row")
)
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
proxy_logging.db_spend_update_writer = MagicMock()
proxy_logging.db_spend_update_writer.spend_log_redis_buffer = MagicMock()
with pytest.raises(ValueError, match="invalid spend log row"):
await ProxyUpdateSpend.update_spend_logs(
n_retry_times=0,
prisma_client=mock_prisma_client,
db_writer_client=None,
proxy_logging_obj=proxy_logging,
logs_to_process=failed_logs,
)
assert mock_prisma_client.spend_log_transactions == []
@pytest.mark.asyncio
async def test_update_spend_logs_retries_deadlock_then_succeeds_without_requeue(
mock_prisma_client: Any,
make_spend_log_row: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.utils as utils_mod
async def _fake_sleep(_: float) -> None:
return None
monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep)
create_many = AsyncMock(
side_effect=[
prisma.errors.PrismaError("deadlock detected"),
None,
]
)
mock_prisma_client.db.litellm_spendlogs.create_many = create_many
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
proxy_logging.db_spend_update_writer = MagicMock()
proxy_logging.db_spend_update_writer.spend_log_redis_buffer = MagicMock()
proxy_logging.db_spend_update_writer.spend_log_redis_buffer.requeue_spend_log_rows = AsyncMock()
logs = [make_spend_log_row(request_id="r1")]
await ProxyUpdateSpend.update_spend_logs(
n_retry_times=1,
prisma_client=mock_prisma_client,
db_writer_client=None,
proxy_logging_obj=proxy_logging,
logs_to_process=logs,
)
assert create_many.await_count == 2
assert mock_prisma_client.spend_log_transactions == []
proxy_logging.db_spend_update_writer.spend_log_redis_buffer.requeue_spend_log_rows.assert_not_awaited()

View file

@ -1,103 +0,0 @@
"""Tests for spend-log flush retry configuration."""
from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.constants import (
DEFAULT_SPEND_LOG_FLUSH_MAX_RETRIES,
SPEND_LOG_FLUSH_MAX_RETRIES,
)
from litellm.proxy.utils import get_spend_log_flush_max_retries, update_spend_logs_job
def test_get_spend_log_flush_max_retries_defaults_to_constant(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(proxy_server_mod, "general_settings", {})
assert get_spend_log_flush_max_retries() == SPEND_LOG_FLUSH_MAX_RETRIES
def test_get_spend_log_flush_max_retries_general_settings_overrides_env(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.constants as constants_mod
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(constants_mod, "SPEND_LOG_FLUSH_MAX_RETRIES", 5)
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"spend_log_flush_max_retries": 7},
)
assert get_spend_log_flush_max_retries() == 7
def test_get_spend_log_flush_max_retries_clamps_negative_to_zero(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(
proxy_server_mod,
"general_settings",
{"spend_log_flush_max_retries": -2},
)
assert get_spend_log_flush_max_retries() == 0
def test_get_spend_log_flush_max_retries_env_default(
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.constants as constants_mod
import litellm.proxy.proxy_server as proxy_server_mod
monkeypatch.setattr(constants_mod, "SPEND_LOG_FLUSH_MAX_RETRIES", 5)
monkeypatch.setattr(proxy_server_mod, "general_settings", {})
assert get_spend_log_flush_max_retries() == 5
@pytest.mark.asyncio
async def test_update_spend_logs_job_uses_configured_max_retries(
mock_prisma_client: object,
make_spend_log_row: object,
monkeypatch: pytest.MonkeyPatch,
) -> None:
import litellm.proxy.guardrails.usage_tracking as guard_mod
import litellm.proxy.db.spend_log_tool_index as tool_mod
import litellm.proxy.utils as utils_mod
monkeypatch.setattr(utils_mod, "get_spend_log_flush_max_retries", lambda: 2)
captured: dict[str, int] = {}
async def _capture_update_spend_logs(**kwargs: object) -> None:
captured["n_retry_times"] = kwargs["n_retry_times"] # type: ignore[index]
monkeypatch.setattr(
utils_mod.ProxyUpdateSpend,
"update_spend_logs",
staticmethod(_capture_update_spend_logs),
)
monkeypatch.setattr(
guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False
)
monkeypatch.setattr(
tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False
)
mock_prisma_client.spend_log_transactions = [ # type: ignore[attr-defined]
make_spend_log_row(request_id="r1") # type: ignore[operator]
]
proxy_logging = MagicMock()
await update_spend_logs_job(
prisma_client=mock_prisma_client, # type: ignore[arg-type]
db_writer_client=None,
proxy_logging_obj=proxy_logging,
)
assert captured["n_retry_times"] == 2