mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): book idempotent /spend/usage records in one transaction, report disabled spend updates
This commit is contained in:
parent
1fc0abe027
commit
73f31beb0d
2 changed files with 88 additions and 44 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue