From 73f31beb0d6b43777b1a05d01e7c039f6d178e5d Mon Sep 17 00:00:00 2001 From: todayim <809634488@qq.com> Date: Thu, 6 Aug 2026 01:00:04 +0800 Subject: [PATCH] fix(proxy): book idempotent /spend/usage records in one transaction, report disabled spend updates --- .../usage_ingestion_endpoints.py | 102 ++++++++++-------- .../test_usage_ingestion_endpoints.py | 30 +++++- 2 files changed, 88 insertions(+), 44 deletions(-) diff --git a/litellm/proxy/spend_tracking/usage_ingestion_endpoints.py b/litellm/proxy/spend_tracking/usage_ingestion_endpoints.py index 2af7f97bb97..9a49268fdae 100644 --- a/litellm/proxy/spend_tracking/usage_ingestion_endpoints.py +++ b/litellm/proxy/spend_tracking/usage_ingestion_endpoints.py @@ -4,7 +4,15 @@ from collections.abc import Awaitable, Callable, Mapping from dataclasses import dataclass from datetime import datetime from types import MappingProxyType -from typing import Annotated, Final, Literal, NamedTuple +from typing import ( + TYPE_CHECKING, + Annotated, + Final, + Literal, + NamedTuple, + TypeAlias, + cast, # noqa: TID251 # untyped tx boundary needs cast for the shim +) from fastapi import APIRouter, Depends, HTTPException, status from pydantic import BaseModel, Field, model_validator @@ -65,10 +73,13 @@ class KeyAttribution(NamedTuple): organization_id: str | None +ReservationOutcome: TypeAlias = Literal["reserved", "duplicate", "disabled"] + + @dataclass(frozen=True, slots=True) class UsageIngestionDeps: lookup_key: Callable[[str], Awaitable[KeyAttribution | None]] - reserve_spend_log: Callable[[ExternalUsageRecord, str, str, KeyAttribution, float], Awaitable[bool]] + reserve_spend_log: Callable[[ExternalUsageRecord, str, str, KeyAttribution, float], Awaitable[ReservationOutcome]] record_spend: Callable[..., Awaitable[None]] compute_cost: Callable[[litellm.ModelResponse, str], float] generate_request_id: Callable[[], str] @@ -161,9 +172,23 @@ async def process_external_usage_record( ) if record.idempotency_key is not None: - reserved: Final = await deps.reserve_spend_log(record, request_id, hashed_token, key, cost) - if reserved is False: + try: + reservation: Final = await deps.reserve_spend_log(record, request_id, hashed_token, key, cost) + except Exception as e: # noqa: BLE001 # booking raises arbitrary persistence errors; an aborted transaction means nothing was booked, so telling the caller to retry is safe + verbose_proxy_logger.info("ingest usage: transactional booking failed for %s: %s", request_id, e) + return UsageIngestRecordResult( + request_id=request_id, + status="error", + error="booking failed transactionally, nothing was recorded, safe to retry", + ) + if reservation == "duplicate": return UsageIngestRecordResult(request_id=request_id, status="duplicate") + if reservation == "disabled": + return UsageIngestRecordResult( + request_id=request_id, + status="error", + error="spend updates are disabled on this proxy, nothing was recorded", + ) return UsageIngestRecordResult(request_id=request_id, status="recorded", spend=cost) await deps.record_spend( @@ -189,60 +214,53 @@ def _attribution_of(key_row: object) -> KeyAttribution: ) +if TYPE_CHECKING: + from prisma.client import TransactionManager + + +class _TransactionClientShim: + def __init__(self, tx: "TransactionManager") -> None: + self.db: Final = tx + + async def reserve_spend_log_atomic( record: ExternalUsageRecord, request_id: str, hashed_token: str, key: KeyAttribution, cost: float, -) -> bool: +) -> ReservationOutcome: from litellm.proxy.proxy_server import litellm_proxy_budget_name, prisma_client, proxy_logging_obj - from litellm.proxy.utils import ProxyUpdateSpend + from litellm.proxy.utils import PrismaClient, ProxyUpdateSpend from litellm.repositories.table_repositories import SpendLogsRepository if ProxyUpdateSpend.disable_spend_updates() is True: - return True + return "disabled" payload: Final = prisma_client.jsonify_object(build_spend_log_payload(record, request_id, hashed_token, key, cost)) from prisma.errors import UniqueViolationError - try: - await SpendLogsRepository(prisma_client).table.create(data=payload) - except UniqueViolationError: - return False - writer: Final = proxy_logging_obj.db_spend_update_writer - counter_calls: Final = ( - writer._update_key_db( - response_cost=cost, - hashed_token=hashed_token, - prisma_client=prisma_client, - ), - writer._update_user_db( - response_cost=cost, - user_id=key.user_id, - prisma_client=prisma_client, - litellm_proxy_budget_name=litellm_proxy_budget_name, - end_user_id=record.end_user_id, - ), - writer._update_team_db( - response_cost=cost, - team_id=key.team_id, - user_id=key.user_id, - prisma_client=prisma_client, - ), - writer._update_org_db( - response_cost=cost, - org_id=key.organization_id, - prisma_client=prisma_client, - ), - ) - results: Final = await asyncio.gather(*counter_calls, return_exceptions=True) - for counter_result in results: - if isinstance(counter_result, Exception): - verbose_proxy_logger.debug("ingest usage: spend counter update failed: %s", counter_result) - return True + try: + async with prisma_client.tx() as tx: + shim: Final = cast(PrismaClient, _TransactionClientShim(tx)) # cast-ok: helper uses only .db (untyped) + await SpendLogsRepository(shim).table.create(data=payload) + await writer._update_key_db(response_cost=cost, hashed_token=hashed_token, prisma_client=shim) + await writer._update_user_db( + response_cost=cost, + user_id=key.user_id, + prisma_client=shim, + litellm_proxy_budget_name=litellm_proxy_budget_name, + end_user_id=record.end_user_id, + ) + await writer._update_team_db( + response_cost=cost, team_id=key.team_id, user_id=key.user_id, prisma_client=shim + ) + await writer._update_org_db(response_cost=cost, org_id=key.organization_id, prisma_client=shim) + except UniqueViolationError: + return "duplicate" + return "reserved" def default_ingestion_deps() -> UsageIngestionDeps: diff --git a/tests/test_litellm/proxy/spend_tracking/test_usage_ingestion_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_usage_ingestion_endpoints.py index 5384cd18a82..b5ae287269a 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_usage_ingestion_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_usage_ingestion_endpoints.py @@ -31,11 +31,15 @@ class RecordingDeps: existing_ids: frozenset[str] = frozenset(), compute_cost_result: float = 0.05, compute_cost_error: Exception | None = None, + reserve_outcome: str = "reserved", + reserve_raises: Exception | None = None, ): self._key = key self._existing_ids = existing_ids self._compute_cost_result = compute_cost_result self._compute_cost_error = compute_cost_error + self._reserve_outcome = reserve_outcome + self._reserve_raises = reserve_raises self.spend_calls: list[dict[str, Any]] = [] self.reserve_calls: list[dict[str, Any]] = [] self.compute_cost_calls: list[tuple[Any, str]] = [] @@ -51,7 +55,7 @@ class RecordingDeps: hashed_token: str, key: KeyAttribution, cost: float, - ) -> bool: + ) -> Any: self.reserve_calls.append( { "record": record, @@ -61,7 +65,11 @@ class RecordingDeps: "cost": cost, } ) - return request_id not in self._existing_ids + if self._reserve_raises is not None: + raise self._reserve_raises + if self._reserve_outcome == "disabled": + return "disabled" + return "duplicate" if request_id in self._existing_ids else "reserved" async def record_spend(**kwargs: Any) -> None: self.spend_calls.append(kwargs) @@ -230,3 +238,21 @@ def test_record_without_idempotency_key_still_flows_tags_to_funnel_kwargs(): assert metadata["tags"] == ["batch:job-42"] assert metadata["user_api_key_end_user_id"] == "tenant-a" assert deps.spend_calls[0]["end_user_id"] == "tenant-a" + + +def test_disabled_spend_updates_reports_error_instead_of_fake_recorded(): + deps = RecordingDeps(reserve_outcome="disabled") + result = run(process_external_usage_record(make_record(cost=0.01, idempotency_key="k-9"), deps.as_deps())) + assert result.status == "error" + assert "disabled" in (result.error or "") + assert result.spend is None + assert deps.spend_calls == [] + + +def test_failed_booking_is_retry_safe_error_not_permanent_duplicate(): + deps = RecordingDeps(reserve_raises=RuntimeError("db gone mid-tx")) + result = run(process_external_usage_record(make_record(cost=0.01, idempotency_key="k-10"), deps.as_deps())) + assert result.status == "error" + assert "safe to retry" in (result.error or "") + assert result.spend is None + assert deps.spend_calls == []