From a2068820eeac3cd9936d4fcced2e70af10dbf4d8 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:07:18 -0700 Subject: [PATCH] perf(proxy): batch daily model usage writes instead of upserting per request (#44243) * perf(proxy): aggregate daily model usage per flush instead of upserting per request * perf(proxy): drain queued model usage in the spend log flush job * perf(proxy): queue model usage at request time instead of writing to the db * test(proxy): cover batched daily model usage aggregation and retries * test(proxy): read back model insights written by the batched flush * fix(proxy): drain the whole model usage queue each flush so it cannot grow unbounded * test(proxy): give the mock prisma client a model usage queue * test(proxy): cover draining a model usage queue larger than one spend log batch --- litellm/proxy/db/db_spend_update_writer.py | 25 +- litellm/proxy/db/model_usage_rollup.py | 165 +++++++++--- litellm/proxy/utils.py | 20 ++ .../unit/proxy/db/test_model_usage_rollup.py | 248 +++++++++++++----- .../test_model_insights_endpoints.py | 23 +- tests/unit/proxy/test_update_spend.py | 2 + .../proxy/utils/prisma_and_spend/conftest.py | 2 + .../prisma_and_spend/test_spend_functions.py | 27 ++ 8 files changed, 392 insertions(+), 120 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ef5dd3f663c..c64ef72ace6 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -70,6 +70,7 @@ from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( WindowSpendUpdateQueue, ) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler +from litellm.proxy.db.model_usage_rollup import build_model_usage_transaction from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING from litellm.proxy.spend_tracking.compression_savings import ( extract_compression_saved_tokens, @@ -728,6 +729,20 @@ class DBSpendUpdateWriter: except Exception as e: verbose_proxy_logger.debug("_enqueue_tool_usage_transaction error (non-blocking): %s", e) + async def _enqueue_model_usage_transaction( + self, + payload: SpendLogsPayload, + prisma_client: PrismaClient, + ) -> None: + try: + transaction: Final = build_model_usage_transaction(payload) + if transaction is None: + return + async with prisma_client._model_usage_transactions_lock: + prisma_client.model_usage_transactions.append(transaction) + except Exception as e: + verbose_proxy_logger.debug("_enqueue_model_usage_transaction error (non-blocking): %s", e) + async def _enqueue_autorouter_turn_transaction( self, payload: SpendLogsPayload, @@ -1144,15 +1159,7 @@ class DBSpendUpdateWriter: traceback.format_exc(), ) - try: - from litellm.proxy.db.model_usage_rollup import increment_daily_model_usage - - await increment_daily_model_usage(prisma_client=prisma_client, payload=payload_copy) - except Exception: - verbose_proxy_logger.debug( - "_batch_database_updates: increment_daily_model_usage failed: %s", - traceback.format_exc(), - ) + await self._enqueue_model_usage_transaction(payload=payload_copy, prisma_client=prisma_client) async def _update_key_db( self, diff --git a/litellm/proxy/db/model_usage_rollup.py b/litellm/proxy/db/model_usage_rollup.py index acd9130da30..98d53573f30 100644 --- a/litellm/proxy/db/model_usage_rollup.py +++ b/litellm/proxy/db/model_usage_rollup.py @@ -1,5 +1,12 @@ +from __future__ import annotations + +import asyncio +import random +from collections.abc import Mapping, Sequence +from dataclasses import dataclass from datetime import datetime -from typing import Final +from itertools import groupby +from typing import TYPE_CHECKING, Final, Protocol from pydantic import TypeAdapter, ValidationError @@ -8,15 +15,48 @@ from litellm.constants import ( MODEL_INSIGHTS_DEFAULT_TASK, MODEL_INSIGHTS_TASK_TAG_PREFIX, ) -from litellm.proxy._types import SpendLogsPayload +from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES, SpendLogsPayload from litellm.proxy.db.model_insights_tasks import load_model_insight_tasks -from litellm.proxy.utils import PrismaClient -from litellm.repositories.table_repositories import DailyModelUsageRepository + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient _METADATA: Final = TypeAdapter(dict[str, object]) _TAGS: Final = TypeAdapter(list[object]) +class _UpsertTable(Protocol): + def upsert(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... + + +class _ModelUsageBatch(Protocol): + litellm_dailymodelusage: _UpsertTable + + +class _ModelUsageBatchManager(Protocol): + async def __aenter__(self) -> _ModelUsageBatch: ... + + async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... + + +@dataclass(frozen=True, slots=True) +class ModelUsageKey: + date: str + model_group: str + model: str + custom_llm_provider: str + task_type: str + + +@dataclass(frozen=True, slots=True) +class ModelUsageTransaction: + key: ModelUsageKey + spend: float + prompt_tokens: int + completion_tokens: int + successful: bool + + def model_usage_task_type(request_tags: str) -> str: try: tags: Final = _TAGS.validate_json(request_tags) @@ -48,43 +88,88 @@ def _date_from_start_time(start_time: datetime | str) -> str | None: return start_time[:10] if len(start_time) >= 10 else None -async def increment_daily_model_usage(prisma_client: PrismaClient, payload: SpendLogsPayload) -> None: +def build_model_usage_transaction(payload: SpendLogsPayload) -> ModelUsageTransaction | None: date: Final = _date_from_start_time(payload["startTime"]) if date is None or _is_internal_call(payload["metadata"]): - return - + return None model: Final = payload["model"] or "unknown" - model_group: Final = payload["model_group"] or model - provider: Final = payload["custom_llm_provider"] or "unknown" - task_type: Final = model_usage_task_type(payload["request_tags"]) - successful: Final = 1 if payload["status"] == "success" else 0 - failed: Final = 1 - successful - key: Final = { - "date": date, - "model_group": model_group, - "model": model, - "custom_llm_provider": provider, - "task_type": task_type, - } - await DailyModelUsageRepository(prisma_client).table.upsert( - where={"date_model_group_model_custom_llm_provider_task_type": key}, - data={ - "create": { - **key, - "spend": payload["spend"], - "prompt_tokens": payload["prompt_tokens"], - "completion_tokens": payload["completion_tokens"], - "request_count": 1, - "successful_requests": successful, - "failed_requests": failed, - }, - "update": { - "spend": {"increment": payload["spend"]}, - "prompt_tokens": {"increment": payload["prompt_tokens"]}, - "completion_tokens": {"increment": payload["completion_tokens"]}, - "request_count": {"increment": 1}, - "successful_requests": {"increment": successful}, - "failed_requests": {"increment": failed}, - }, - }, + return ModelUsageTransaction( + key=ModelUsageKey( + date=date, + model_group=payload["model_group"] or model, + model=model, + custom_llm_provider=payload["custom_llm_provider"] or "unknown", + task_type=model_usage_task_type(payload["request_tags"]), + ), + spend=payload["spend"], + prompt_tokens=payload["prompt_tokens"], + completion_tokens=payload["completion_tokens"], + successful=payload["status"] == "success", ) + + +def _model_usage_batch(prisma_client: PrismaClient) -> _ModelUsageBatchManager: + batch: Final[_ModelUsageBatchManager] = prisma_client.db.batch_() + return batch + + +def _sort_key(transaction: ModelUsageTransaction) -> tuple[str, str, str, str, str]: + key: Final = transaction.key + return (key.date, key.model_group, key.model, key.custom_llm_provider, key.task_type) + + +async def flush_model_usage_transactions( + prisma_client: PrismaClient, + transactions: Sequence[ModelUsageTransaction], + n_retry_times: int = 3, +) -> None: + """One upsert per rollup row for the whole drained batch, in a single transaction and in key order so + concurrent pods take row locks in the same order. Only ConnectError is retried: it proves nothing reached + the database, while a retry after an ambiguous post-send failure could double-count the increments.""" + if not transactions: + return + ordered: Final = sorted(transactions, key=_sort_key) + for attempt in range(n_retry_times + 1): + try: + async with _model_usage_batch(prisma_client) as batcher: + for key, grouped in groupby(ordered, key=lambda transaction: transaction.key): + entries = tuple(grouped) + spend = sum(entry.spend for entry in entries) + prompt_tokens = sum(entry.prompt_tokens for entry in entries) + completion_tokens = sum(entry.completion_tokens for entry in entries) + successful = sum(1 for entry in entries if entry.successful) + failed = len(entries) - successful + key_fields = { + "date": key.date, + "model_group": key.model_group, + "model": key.model, + "custom_llm_provider": key.custom_llm_provider, + "task_type": key.task_type, + } + batcher.litellm_dailymodelusage.upsert( + where={"date_model_group_model_custom_llm_provider_task_type": key_fields}, + data={ + "create": { + **key_fields, + "spend": spend, + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "request_count": len(entries), + "successful_requests": successful, + "failed_requests": failed, + }, + "update": { + "spend": {"increment": spend}, + "prompt_tokens": {"increment": prompt_tokens}, + "completion_tokens": {"increment": completion_tokens}, + "request_count": {"increment": len(entries)}, + "successful_requests": {"increment": successful}, + "failed_requests": {"increment": failed}, + }, + }, + ) + return + except DB_RETRY_SAFE_ERROR_TYPES: + if attempt >= n_retry_times: + raise + await asyncio.sleep(2.0**attempt + random.uniform(0, 1)) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 817231bc0b9..5828ed94984 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -266,6 +266,7 @@ if TYPE_CHECKING: from litellm.models.team import LiteLLM_TeamTableCachedObj from litellm.proxy.db.autorouter_session_rollup import AutoRouterTurnTransaction from litellm.proxy.db.baseline_accounting import BaselineAccountingRecord + from litellm.proxy.db.model_usage_rollup import ModelUsageTransaction from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline @@ -4400,6 +4401,8 @@ class PrismaClient: spend_log_write_lock = asyncio.Lock() tool_usage_transactions: list["ToolUsageTransaction"] = [] _tool_usage_transactions_lock = asyncio.Lock() + model_usage_transactions: ClassVar[list["ModelUsageTransaction"]] = [] + _model_usage_transactions_lock = asyncio.Lock() autorouter_turn_transactions: ClassVar[list["AutoRouterTurnTransaction"]] = [] _autorouter_turn_transactions_lock = asyncio.Lock() @@ -7508,6 +7511,8 @@ async def _total_queued_spend_transactions(prisma_client: PrismaClient) -> int: spend_queue_size: Final = len(prisma_client.spend_log_transactions) async with prisma_client._tool_usage_transactions_lock: tool_queue_size: Final = len(prisma_client.tool_usage_transactions) + async with prisma_client._model_usage_transactions_lock: + model_usage_queue_size: Final = len(prisma_client.model_usage_transactions) async with prisma_client._autorouter_turn_transactions_lock: autorouter_queue_size: Final = len(prisma_client.autorouter_turn_transactions) from litellm.proxy.db.shadow_eval_funnel import pending_shadow_eval_funnel_events @@ -7517,6 +7522,7 @@ async def _total_queued_spend_transactions(prisma_client: PrismaClient) -> int: return ( spend_queue_size + tool_queue_size + + model_usage_queue_size + autorouter_queue_size + baseline_queue_size + pending_shadow_eval_funnel_events() @@ -7652,6 +7658,20 @@ async def _run_spend_logs_job( tool_tracking_err, ) + async with prisma_client._model_usage_transactions_lock: + model_usage_to_process: Final = prisma_client.model_usage_transactions + prisma_client.model_usage_transactions = [] + try: + from litellm.proxy.db.model_usage_rollup import flush_model_usage_transactions + + await flush_model_usage_transactions(prisma_client=prisma_client, transactions=model_usage_to_process) + except Exception as model_usage_err: + verbose_proxy_logger.error( + "Spend tracking - model usage flush failed; %s model usage transactions dropped: %s", + len(model_usage_to_process), + model_usage_err, + ) + await flush_baseline_accounting(prisma_client) async with prisma_client._autorouter_turn_transactions_lock: diff --git a/tests/unit/proxy/db/test_model_usage_rollup.py b/tests/unit/proxy/db/test_model_usage_rollup.py index f54856129dc..b9b806f4c54 100644 --- a/tests/unit/proxy/db/test_model_usage_rollup.py +++ b/tests/unit/proxy/db/test_model_usage_rollup.py @@ -1,9 +1,68 @@ +import asyncio from datetime import datetime, timezone -from unittest.mock import AsyncMock, MagicMock +from typing import Any +from unittest.mock import MagicMock +import httpx import pytest -from litellm.proxy.db.model_usage_rollup import increment_daily_model_usage, model_usage_task_type +from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter +from litellm.proxy.db.model_usage_rollup import ( + ModelUsageKey, + ModelUsageTransaction, + build_model_usage_transaction, + flush_model_usage_transactions, + model_usage_task_type, +) + + +class _FakeBatcher: + def __init__(self) -> None: + self.litellm_dailymodelusage = MagicMock() + + async def __aenter__(self) -> "_FakeBatcher": + return self + + async def __aexit__(self, *args: Any) -> None: + return None + + +def _prisma(batch_: MagicMock) -> MagicMock: + prisma = MagicMock() + prisma.db.batch_ = batch_ + return prisma + + +def _payload(**overrides: Any) -> dict[str, Any]: + return { + "spend": 0.25, + "prompt_tokens": 10, + "completion_tokens": 20, + "startTime": datetime(2026, 9, 28, 13, tzinfo=timezone.utc), + "model": "openai/gpt-5.4-mini", + "model_group": "fast-chat", + "metadata": "{}", + "request_tags": "[]", + "custom_llm_provider": "openai", + "status": "success", + **overrides, + } + + +def _transaction(model: str, spend: float, successful: bool = True) -> ModelUsageTransaction: + return ModelUsageTransaction( + key=ModelUsageKey( + date="2026-09-28", model_group=model, model=model, custom_llm_provider="openai", task_type="debugging" + ), + spend=spend, + prompt_tokens=10, + completion_tokens=5, + successful=successful, + ) + + +async def _no_sleep(seconds: float) -> None: + return None def test_model_usage_task_type_reads_task_tag_or_defaults() -> None: @@ -14,76 +73,131 @@ def test_model_usage_task_type_reads_task_tag_or_defaults() -> None: assert model_usage_task_type("not json") == "uncategorized" -@pytest.mark.asyncio -async def test_increment_daily_model_usage_uses_atomic_prisma_upsert() -> None: - table = MagicMock() - table.upsert = AsyncMock() - prisma_client = MagicMock() - prisma_client.db.litellm_dailymodelusage = table - payload = { - "request_id": "request-1", - "call_type": "acompletion", - "api_key": "key", - "spend": 0.25, - "total_tokens": 30, - "prompt_tokens": 10, - "completion_tokens": 20, - "startTime": datetime(2026, 9, 28, tzinfo=timezone.utc), - "endTime": datetime(2026, 9, 28, tzinfo=timezone.utc), - "completionStartTime": None, - "model": "openai/gpt-5.4-mini", - "model_id": None, - "model_group": "fast-chat", - "mcp_namespaced_tool_name": None, - "agent_id": None, - "api_base": "", - "user": "user", - "metadata": "{}", - "cache_hit": "False", - "cache_key": "", - "request_tags": "[]", - "team_id": None, - "organization_id": None, - "end_user": None, - "requester_ip_address": None, - "custom_llm_provider": "openai", - "messages": None, - "response": None, - "proxy_server_request": None, - "session_id": None, - "request_duration_ms": 20, - "status": "success", - "litellm_call_id": None, - } +def test_build_model_usage_transaction_keys_on_day_model_and_task() -> None: + transaction = build_model_usage_transaction(_payload(request_tags='["task:debugging"]', status="failure")) - await increment_daily_model_usage(prisma_client, payload) + assert transaction == ModelUsageTransaction( + key=ModelUsageKey( + date="2026-09-28", + model_group="fast-chat", + model="openai/gpt-5.4-mini", + custom_llm_provider="openai", + task_type="debugging", + ), + spend=0.25, + prompt_tokens=10, + completion_tokens=20, + successful=False, + ) - call = table.upsert.await_args.kwargs - assert call["data"]["create"]["request_count"] == 1 - assert call["data"]["update"]["completion_tokens"] == {"increment": 20} - assert call["data"]["create"]["task_type"] == "uncategorized" + +def test_build_model_usage_transaction_falls_back_for_missing_model_fields() -> None: + transaction = build_model_usage_transaction( + _payload(model="", model_group=None, custom_llm_provider=None, startTime="2026-09-28T01:02:03Z") + ) + + assert transaction is not None + assert transaction.key == ModelUsageKey( + date="2026-09-28", + model_group="unknown", + model="unknown", + custom_llm_provider="unknown", + task_type="uncategorized", + ) + + +@pytest.mark.parametrize( + "overrides", + [{"metadata": '{"internal_call_origin": "health_check"}'}, {"startTime": "bad"}], +) +def test_build_model_usage_transaction_skips_internal_calls_and_bad_dates(overrides: dict[str, Any]) -> None: + assert build_model_usage_transaction(_payload(**overrides)) is None @pytest.mark.asyncio -async def test_increment_daily_model_usage_records_task_from_request_tags() -> None: - table = MagicMock() - table.upsert = AsyncMock() - prisma_client = MagicMock() - prisma_client.db.litellm_dailymodelusage = table - payload = { - "call_type": "acompletion", - "spend": 0.1, - "prompt_tokens": 1, - "completion_tokens": 2, - "startTime": datetime(2026, 9, 28, tzinfo=timezone.utc), - "model": "gpt-5", - "model_group": "gpt-5", - "metadata": "{}", - "request_tags": '["task:debugging"]', - "custom_llm_provider": "openai", - "status": "success", +async def test_flush_aggregates_each_rollup_row_into_one_upsert() -> None: + batcher = _FakeBatcher() + prisma = _prisma(MagicMock(return_value=batcher)) + + await flush_model_usage_transactions( + prisma_client=prisma, + transactions=[ + _transaction("gpt-5", 0.5), + _transaction("claude", 1.0), + _transaction("gpt-5", 0.25, successful=False), + _transaction("gpt-5", 0.25), + ], + ) + + upserts = { + call.kwargs["where"]["date_model_group_model_custom_llm_provider_task_type"]["model"]: call.kwargs["data"] + for call in batcher.litellm_dailymodelusage.upsert.call_args_list } + assert list(upserts) == ["claude", "gpt-5"] + gpt = upserts["gpt-5"] + assert gpt["create"]["spend"] == 1.0 + assert gpt["create"]["prompt_tokens"] == 30 + assert gpt["create"]["completion_tokens"] == 15 + assert gpt["create"]["request_count"] == 3 + assert gpt["create"]["successful_requests"] == 2 + assert gpt["create"]["failed_requests"] == 1 + assert gpt["update"] == { + "spend": {"increment": 1.0}, + "prompt_tokens": {"increment": 30}, + "completion_tokens": {"increment": 15}, + "request_count": {"increment": 3}, + "successful_requests": {"increment": 2}, + "failed_requests": {"increment": 1}, + } + assert upserts["claude"]["create"]["request_count"] == 1 - await increment_daily_model_usage(prisma_client, payload) - assert table.upsert.await_args.kwargs["data"]["create"]["task_type"] == "debugging" +@pytest.mark.asyncio +async def test_flush_with_no_transactions_touches_nothing() -> None: + prisma = _prisma(MagicMock()) + await flush_model_usage_transactions(prisma_client=prisma, transactions=[]) + prisma.db.batch_.assert_not_called() + + +@pytest.mark.asyncio +async def test_flush_retries_connection_errors(monkeypatch: pytest.MonkeyPatch) -> None: + batcher = _FakeBatcher() + prisma = _prisma(MagicMock(side_effect=[httpx.ConnectError("down"), batcher])) + monkeypatch.setattr("litellm.proxy.db.model_usage_rollup.asyncio.sleep", _no_sleep) + + await flush_model_usage_transactions(prisma_client=prisma, transactions=[_transaction("gpt-5", 0.1)]) + + assert prisma.db.batch_.call_count == 2 + batcher.litellm_dailymodelusage.upsert.assert_called_once() + + +@pytest.mark.asyncio +async def test_flush_does_not_retry_ambiguous_errors() -> None: + prisma = _prisma(MagicMock(side_effect=httpx.ReadTimeout("ambiguous"))) + + with pytest.raises(httpx.ReadTimeout): + await flush_model_usage_transactions(prisma_client=prisma, transactions=[_transaction("gpt-5", 0.1)]) + + prisma.db.batch_.assert_called_once() + + +@pytest.mark.asyncio +async def test_request_time_path_queues_usage_instead_of_writing_to_the_db() -> None: + prisma = MagicMock() + prisma.model_usage_transactions = [] + prisma._model_usage_transactions_lock = asyncio.Lock() + + await DBSpendUpdateWriter()._batch_database_updates( + response_cost=0.25, + user_id="u1", + hashed_token="t1", + team_id=None, + org_id=None, + end_user_id=None, + prisma_client=prisma, + litellm_proxy_budget_name=None, + payload=_payload(request_id="req-1"), + ) + + assert [transaction.key.model for transaction in prisma.model_usage_transactions] == ["openai/gpt-5.4-mini"] + prisma.db.litellm_dailymodelusage.upsert.assert_not_called() diff --git a/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py index 535f32a7f10..7c8d1346944 100644 --- a/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_insights_endpoints.py @@ -7,7 +7,7 @@ from fastapi.testclient import TestClient from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.db.model_usage_rollup import increment_daily_model_usage +from litellm.proxy.db.model_usage_rollup import build_model_usage_transaction, flush_model_usage_transactions from litellm.proxy.management_endpoints.model_insights_endpoints import router @@ -199,7 +199,7 @@ class _InMemoryUsageTable: def __init__(self) -> None: self.rows: dict[tuple[str, ...], dict[str, float]] = {} - async def upsert(self, where: dict, data: dict) -> None: + def upsert(self, where: dict, data: dict) -> None: key_fields = where["date_model_group_model_custom_llm_provider_task_type"] key = tuple(key_fields.values()) if key not in self.rows: @@ -221,11 +221,23 @@ class _InMemoryUsageTable: return list(grouped.values()) +class _InMemoryBatcher: + def __init__(self, table: _InMemoryUsageTable) -> None: + self.litellm_dailymodelusage = table + + async def __aenter__(self) -> "_InMemoryBatcher": + return self + + async def __aexit__(self, *args: object) -> None: + return None + + @pytest.mark.asyncio async def test_model_insights_reads_back_what_the_rollup_wrote() -> None: table = _InMemoryUsageTable() prisma = MagicMock() prisma.db.litellm_dailymodelusage = table + prisma.db.batch_ = MagicMock(return_value=_InMemoryBatcher(table)) payload = { "call_type": "acompletion", "spend": 0.5, @@ -240,8 +252,11 @@ async def test_model_insights_reads_back_what_the_rollup_wrote() -> None: "status": "success", } - await increment_daily_model_usage(prisma, payload) - await increment_daily_model_usage(prisma, {**payload, "request_tags": "[]"}) + transactions = ( + build_model_usage_transaction(payload), + build_model_usage_transaction({**payload, "request_tags": "[]"}), + ) + await flush_model_usage_transactions(prisma, [t for t in transactions if t is not None]) body = _call(table, "metric=requests").json() diff --git a/tests/unit/proxy/test_update_spend.py b/tests/unit/proxy/test_update_spend.py index ebe505b3d60..6b92320762b 100644 --- a/tests/unit/proxy/test_update_spend.py +++ b/tests/unit/proxy/test_update_spend.py @@ -36,6 +36,7 @@ class MockPrismaClient: self.spend_log_transactions = [] self.daily_user_spend_transactions = {} self.tool_usage_transactions = [] + self.model_usage_transactions = [] self.autorouter_turn_transactions = [] self.baseline_accounting_transactions = [] self.baseline_accounting_lock = asyncio.Lock() @@ -49,6 +50,7 @@ class MockPrismaClient: self._spend_log_transactions_lock = asyncio.Lock() self.spend_log_write_lock = asyncio.Lock() self._tool_usage_transactions_lock = asyncio.Lock() + self._model_usage_transactions_lock = asyncio.Lock() self._autorouter_turn_transactions_lock = asyncio.Lock() def jsonify_object(self, obj): diff --git a/tests/unit/proxy/utils/prisma_and_spend/conftest.py b/tests/unit/proxy/utils/prisma_and_spend/conftest.py index 455eb423ddc..e37a82a023b 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/conftest.py +++ b/tests/unit/proxy/utils/prisma_and_spend/conftest.py @@ -133,6 +133,8 @@ def mock_prisma_client() -> MagicMock: client.spend_log_write_lock = asyncio.Lock() client.tool_usage_transactions = [] client._tool_usage_transactions_lock = asyncio.Lock() + client.model_usage_transactions = [] + client._model_usage_transactions_lock = asyncio.Lock() client.jsonify_object = lambda data: dict(data) client.db.is_connected = MagicMock(return_value=False) client.db.connect = AsyncMock() diff --git a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py index d6f41ba55db..7aa06bcafae 100644 --- a/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/unit/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -223,6 +223,33 @@ async def test_update_spend_logs_job_drains_tool_queue_when_spend_queue_empty( assert mock_prisma_client.tool_usage_transactions == [] +@pytest.mark.asyncio +async def test_update_spend_logs_job_drains_the_whole_model_usage_queue_in_one_run( + mock_prisma_client: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + import litellm.proxy.db.model_usage_rollup as model_usage_mod + import litellm.proxy.db.spend_log_tool_index as tool_mod + import litellm.proxy.guardrails.usage_tracking as guard_mod + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + queued = [MagicMock() for _ in range(25_000)] + mock_prisma_client.model_usage_transactions = list(queued) + monkeypatch.setattr(guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False) + monkeypatch.setattr(tool_mod, "flush_tool_usage_transactions", AsyncMock(), raising=False) + flush_stub = AsyncMock() + monkeypatch.setattr(model_usage_mod, "flush_model_usage_transactions", flush_stub, raising=False) + + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + assert flush_stub.await_args.kwargs["transactions"] == queued + assert mock_prisma_client.model_usage_transactions == [] + + @pytest.mark.asyncio async def test_update_spend_logs_job_processes_and_clears_queue( mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch