diff --git a/litellm/constants.py b/litellm/constants.py index 36e578bd323..3308702b785 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -295,6 +295,7 @@ 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)) @@ -1494,6 +1495,15 @@ 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) ) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 88be567e59a..47de61d2ada 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2585,6 +2585,10 @@ 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).", @@ -3635,6 +3639,8 @@ 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): @@ -3960,6 +3966,42 @@ 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 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e7f14df5294..2ef4a2a2550 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -55,6 +55,9 @@ 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, @@ -112,6 +115,7 @@ 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() @@ -755,12 +759,12 @@ class DBSpendUpdateWriter: payload.get("request_id"), payload.get("spend") ) ) - 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: + if 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." diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_redis_buffer.py b/litellm/proxy/db/db_transaction_queue/spend_log_redis_buffer.py new file mode 100644 index 00000000000..055b648892b --- /dev/null +++ b/litellm/proxy/db/db_transaction_queue/spend_log_redis_buffer.py @@ -0,0 +1,96 @@ +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 diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 6667010447b..2c3e1ed53e5 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -557,6 +557,8 @@ 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 @@ -1302,6 +1304,7 @@ 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( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e77e24c9e71..d6407233b4c 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -42,6 +42,7 @@ 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 @@ -5340,6 +5341,7 @@ 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): @@ -5361,7 +5363,6 @@ 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): @@ -5376,33 +5377,37 @@ 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 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) + 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 Exception as e: - # Logs already removed from queue at start - don't put them back - # This matches the original behavior where logs are removed even on error + 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, + ) _raise_failed_update_spend_exception( e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) @@ -5516,6 +5521,107 @@ 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], @@ -5527,20 +5633,29 @@ 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. """ - n_retry_times = 3 + # 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() MAX_LOGS_PER_INTERVAL = 10000 - # 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 + 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() async with prisma_client._spend_log_transactions_lock: - 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) : - ] + 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 await ProxyUpdateSpend.update_spend_logs( n_retry_times=n_retry_times, diff --git a/tests/test_litellm/proxy/_types/test_spend_log_flush_retryable_error.py b/tests/test_litellm/proxy/_types/test_spend_log_flush_retryable_error.py new file mode 100644 index 00000000000..43896224407 --- /dev/null +++ b/tests/test_litellm/proxy/_types/test_spend_log_flush_retryable_error.py @@ -0,0 +1,26 @@ +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 diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_log_redis_buffer.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_log_redis_buffer.py new file mode 100644 index 00000000000..2a89b3bd3dd --- /dev/null +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_log_redis_buffer.py @@ -0,0 +1,128 @@ +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 + diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 79e6494eab0..a35a88099c3 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1656,3 +1656,68 @@ 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() + diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_spend_tracking.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_spend_tracking.py new file mode 100644 index 00000000000..26ff911ddc5 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_spend_tracking.py @@ -0,0 +1,350 @@ +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" + diff --git a/tests/test_litellm/proxy/spend_tracking/SPEND_TRACKING_TEST_RUNBOOK.md b/tests/test_litellm/proxy/spend_tracking/SPEND_TRACKING_TEST_RUNBOOK.md deleted file mode 100644 index 3d040ef2b02..00000000000 --- a/tests/test_litellm/proxy/spend_tracking/SPEND_TRACKING_TEST_RUNBOOK.md +++ /dev/null @@ -1,102 +0,0 @@ -# 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 diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py index 6a4fd516c9b..b7bebb63b13 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py @@ -216,6 +216,7 @@ 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") ) @@ -227,8 +228,78 @@ 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=[make_spend_log_row(request_id="r1")], + logs_to_process=failed_logs, ) + 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( diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_durability.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_durability.py new file mode 100644 index 00000000000..905d756748a --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_durability.py @@ -0,0 +1,263 @@ +"""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() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_max_retries.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_max_retries.py new file mode 100644 index 00000000000..47ba043caf9 --- /dev/null +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_max_retries.py @@ -0,0 +1,103 @@ +"""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