mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Revert "fix(proxy): durable spend-log flush with redis buffer and passthrough attribution"
This reverts commit 9a3a16fb53.
This commit is contained in:
parent
9a3a16fb53
commit
586e31f1c4
14 changed files with 136 additions and 1310 deletions
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
@ -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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue