diff --git a/litellm/constants.py b/litellm/constants.py index 3308702b785..36e578bd323 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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) ) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 47de61d2ada..88be567e59a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 2ef4a2a2550..e7f14df5294 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -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." 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 deleted file mode 100644 index 055b648892b..00000000000 --- a/litellm/proxy/db/db_transaction_queue/spend_log_redis_buffer.py +++ /dev/null @@ -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 diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 2c3e1ed53e5..6667010447b 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d6407233b4c..e77e24c9e71 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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, 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 deleted file mode 100644 index 43896224407..00000000000 --- a/tests/test_litellm/proxy/_types/test_spend_log_flush_retryable_error.py +++ /dev/null @@ -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 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 deleted file mode 100644 index 2a89b3bd3dd..00000000000 --- a/tests/test_litellm/proxy/db/db_transaction_queue/test_spend_log_redis_buffer.py +++ /dev/null @@ -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 - 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 a35a88099c3..79e6494eab0 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,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() - 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 deleted file mode 100644 index 26ff911ddc5..00000000000 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_spend_tracking.py +++ /dev/null @@ -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" - 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 new file mode 100644 index 00000000000..3d040ef2b02 --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/SPEND_TRACKING_TEST_RUNBOOK.md @@ -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 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 b7bebb63b13..6a4fd516c9b 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,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( 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 deleted file mode 100644 index 905d756748a..00000000000 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_durability.py +++ /dev/null @@ -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() 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 deleted file mode 100644 index 47ba043caf9..00000000000 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_log_flush_max_retries.py +++ /dev/null @@ -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