mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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
This commit is contained in:
parent
8b4de39ad7
commit
a2068820ee
8 changed files with 392 additions and 120 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue