fix(proxy): book idempotent /spend/usage records in one transaction, report disabled spend updates

This commit is contained in:
todayim 2026-08-06 01:00:04 +08:00
parent 1fc0abe027
commit 73f31beb0d
2 changed files with 88 additions and 44 deletions

View file

@ -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:

View file

@ -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 == []