fix(proxy): enforce proxy-admin role, atomic idempotency, key-helper lookup for /spend/usage

This commit is contained in:
todayim 2026-08-05 17:28:21 +08:00
parent bdfd0e4691
commit b9ce498009
3 changed files with 368 additions and 70 deletions

View file

@ -1,16 +1,18 @@
import asyncio
import uuid
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import datetime
from typing import Final, Literal, NamedTuple
from typing import Annotated, Final, Literal, NamedTuple
from fastapi import APIRouter, Depends, status
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, Field, model_validator
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import ProxyException
from litellm.proxy._types import LitellmUserRoles, ProxyException, SpendLogsPayload, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
from litellm.proxy.utils import hash_token
router: Final = APIRouter()
@ -25,8 +27,12 @@ class ExternalUsageRecord(BaseModel):
completion_tokens: int = Field(ge=0)
start_time: datetime
end_time: datetime | None = None
cost: float | None = Field(default=None, ge=0, description="Explicit cost in USD. Computed from litellm pricing when omitted.")
idempotency_key: str | None = Field(default=None, max_length=255, description="Becomes the spend-log request_id for dedup on retries.")
cost: float | None = Field(
default=None, ge=0, description="Explicit cost in USD. Computed from litellm pricing when omitted."
)
idempotency_key: str | None = Field(
default=None, max_length=255, description="Becomes the spend-log request_id for dedup on retries."
)
tags: list[str] | None = None
end_user_id: str | None = None
@ -58,14 +64,14 @@ class KeyAttribution(NamedTuple):
organization_id: str | None
RecordSpendFn = Callable[..., Awaitable[None]]
ReserveSpendFn = Callable[[ExternalUsageRecord, str, str, KeyAttribution, float], Awaitable[bool]]
@dataclass(frozen=True, slots=True)
class UsageIngestionDeps:
lookup_key: Callable[[str], Awaitable[KeyAttribution | None]]
spend_log_exists: Callable[[str], Awaitable[bool]]
record_spend: RecordSpendFn
reserve_spend_log: ReserveSpendFn
record_spend: Callable[..., Awaitable[None]]
compute_cost: Callable[[litellm.ModelResponse, str], float]
generate_request_id: Callable[[], str]
@ -98,13 +104,40 @@ def _build_completion_response(record: ExternalUsageRecord, request_id: str) ->
)
def build_spend_log_payload(
record: ExternalUsageRecord,
request_id: str,
hashed_token: str,
key: KeyAttribution,
cost: float,
) -> SpendLogsPayload:
payload: SpendLogsPayload = get_logging_payload(
kwargs=_build_usage_kwargs(record, hashed_token),
response_obj=_build_completion_response(record, request_id),
start_time=record.start_time,
end_time=record.end_time or record.start_time,
)
payload["spend"] = cost
if isinstance(payload["startTime"], datetime):
payload["startTime"] = payload["startTime"].isoformat()
if isinstance(payload["endTime"], datetime):
payload["endTime"] = payload["endTime"].isoformat()
if key.organization_id is not None and key.organization_id != "":
payload["organization_id"] = key.organization_id
if key.team_id is not None and key.team_id != "":
payload["team_id"] = key.team_id
return payload
def _resolve_cost(deps: UsageIngestionDeps, record: ExternalUsageRecord, response: litellm.ModelResponse) -> float:
if record.cost is not None:
return record.cost
return deps.compute_cost(response, record.model)
async def process_external_usage_record(record: ExternalUsageRecord, deps: UsageIngestionDeps) -> UsageIngestRecordResult:
async def process_external_usage_record(
record: ExternalUsageRecord, deps: UsageIngestionDeps
) -> UsageIngestRecordResult:
request_id: Final = record.idempotency_key or deps.generate_request_id()
hashed_token: Final = hash_token(record.api_key)
@ -112,14 +145,11 @@ async def process_external_usage_record(record: ExternalUsageRecord, deps: Usage
if key is None:
return UsageIngestRecordResult(request_id=request_id, status="error", error="api key not found")
if record.idempotency_key is not None and await deps.spend_log_exists(request_id):
return UsageIngestRecordResult(request_id=request_id, status="duplicate")
response: Final = _build_completion_response(record, request_id)
try:
cost: Final = _resolve_cost(deps, record, response)
except Exception as e: # noqa: BLE001 # pricing lookup raises arbitrary provider-specific errors; any failure means the record is unpriceable and must carry an explicit cost
except Exception as e: # noqa: BLE001
verbose_proxy_logger.info("ingest usage: cost computation failed for model %s: %s", record.model, e)
return UsageIngestRecordResult(
request_id=request_id,
@ -127,6 +157,12 @@ async def process_external_usage_record(record: ExternalUsageRecord, deps: Usage
error="could not compute cost for this model, pass an explicit cost",
)
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:
return UsageIngestRecordResult(request_id=request_id, status="duplicate")
return UsageIngestRecordResult(request_id=request_id, status="recorded", spend=cost)
await deps.record_spend(
token=hashed_token,
user_id=key.user_id,
@ -150,8 +186,72 @@ def _attribution_of(key_row: object) -> KeyAttribution:
)
async def reserve_spend_log_atomic(
record: ExternalUsageRecord,
request_id: str,
hashed_token: str,
key: KeyAttribution,
cost: float,
) -> bool:
from litellm.proxy.proxy_server import (
disable_spend_logs,
litellm_proxy_budget_name,
prisma_client,
proxy_logging_obj,
)
from litellm.proxy.utils import ProxyUpdateSpend
from litellm.repositories.table_repositories import SpendLogsRepository
if ProxyUpdateSpend.disable_spend_updates() is True:
return True
if disable_spend_logs is False:
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
def default_ingestion_deps() -> UsageIngestionDeps:
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
if prisma_client is None:
raise ProxyException(
@ -162,18 +262,21 @@ def default_ingestion_deps() -> UsageIngestionDeps:
)
async def lookup_key(hashed_token: str) -> KeyAttribution | None:
row: Final = await prisma_client.db.litellm_verificationtoken.find_unique(where={"token": hashed_token})
if row is None:
return None
return _attribution_of(row)
from litellm.proxy.auth.auth_checks import get_key_object
async def spend_log_exists(request_id: str) -> bool:
row: Final = await prisma_client.db.litellm_spendlogs.find_unique(where={"request_id": request_id})
return row is not None
try:
key_row: Final = await get_key_object(
hashed_token=hashed_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
)
except ProxyException:
return None
return _attribution_of(key_row)
return UsageIngestionDeps(
lookup_key=lookup_key,
spend_log_exists=spend_log_exists,
reserve_spend_log=reserve_spend_log_atomic,
record_spend=proxy_logging_obj.db_spend_update_writer.update_database,
compute_cost=lambda resp, model: litellm.completion_cost(completion_response=resp, model=model),
generate_request_id=lambda: str(uuid.uuid4()),
@ -186,7 +289,10 @@ def default_ingestion_deps() -> UsageIngestionDeps:
dependencies=[Depends(user_api_key_auth)],
response_model=UsageIngestResponse,
)
async def ingest_external_usage(request: UsageIngestRequest) -> UsageIngestResponse:
async def ingest_external_usage(
request: UsageIngestRequest,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
) -> UsageIngestResponse:
"""
PROXY_ADMIN ONLY: record externally measured usage into the same spend pipeline as proxy-routed traffic.
@ -195,10 +301,17 @@ async def ingest_external_usage(request: UsageIngestRequest) -> UsageIngestRespo
single metering system.
Attribution (user/team/org) is derived from the given virtual key. Records accept an optional
idempotency_key, stored as the spend-log request_id, so retries are deduplicated. When cost is
omitted it is computed from litellm pricing; records whose model cannot be priced are rejected
with an error instead of being booked as zero spend.
idempotency_key, stored as the spend-log request_id: the reservation insert, counter updates and
dedup are checked atomically at the database primary key, so overlapping retries are safe. When
cost is omitted it is computed from litellm pricing; records whose model cannot be priced are
rejected with an error instead of being booked as zero spend.
"""
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only proxy admins ingest spend records here. Use a key with the proxy_admin role.",
)
deps: Final = default_ingestion_deps()
results: Final = [await process_external_usage_record(record, deps) for record in request.records]
return UsageIngestResponse(results=results)

View file

@ -1,4 +1,5 @@
import asyncio
import json
import os
import sys
from datetime import datetime, timezone
@ -13,6 +14,7 @@ from litellm.proxy.spend_tracking.usage_ingestion_endpoints import (
ExternalUsageRecord,
KeyAttribution,
UsageIngestionDeps,
build_spend_log_payload,
process_external_usage_record,
)
from litellm.proxy.utils import hash_token
@ -35,17 +37,31 @@ class RecordingDeps:
self._compute_cost_result = compute_cost_result
self._compute_cost_error = compute_cost_error
self.spend_calls: list[dict[str, Any]] = []
self.reserve_calls: list[dict[str, Any]] = []
self.compute_cost_calls: list[tuple[Any, str]] = []
self.exists_calls: list[str] = []
def as_deps(self) -> UsageIngestionDeps:
async def lookup_key(hashed: str) -> KeyAttribution | None:
self.looked_up_hashed = hashed
return self._key
async def spend_log_exists(request_id: str) -> bool:
self.exists_calls.append(request_id)
return request_id in self._existing_ids
async def reserve_spend_log(
record: ExternalUsageRecord,
request_id: str,
hashed_token: str,
key: KeyAttribution,
cost: float,
) -> bool:
self.reserve_calls.append(
{
"record": record,
"request_id": request_id,
"hashed_token": hashed_token,
"key": key,
"cost": cost,
}
)
return request_id not in self._existing_ids
async def record_spend(**kwargs: Any) -> None:
self.spend_calls.append(kwargs)
@ -58,7 +74,7 @@ class RecordingDeps:
return UsageIngestionDeps(
lookup_key=lookup_key,
spend_log_exists=spend_log_exists,
reserve_spend_log=reserve_spend_log,
record_spend=record_spend,
compute_cost=compute_cost,
generate_request_id=lambda: GENERATED_ID,
@ -81,61 +97,82 @@ def run(coro: Any) -> Any:
return asyncio.run(coro)
def test_records_spend_with_explicit_cost_without_calling_pricing():
def test_spend_log_payload_matches_funnel_shape():
record = make_record(idempotency_key="batch-1-line-1", tags=["batch:job-42"], end_user_id="tenant-a")
payload = build_spend_log_payload(
record=record,
request_id="batch-1-line-1",
hashed_token=hash_token(RAW_KEY),
key=DEFAULT_KEY,
cost=0.123,
)
assert payload["request_id"] == "batch-1-line-1"
assert payload["spend"] == 0.123
assert payload["total_tokens"] == 150
assert payload["prompt_tokens"] == 100
assert payload["completion_tokens"] == 50
assert payload["api_key"] == hash_token(RAW_KEY)
assert payload["api_key"] != RAW_KEY
assert payload["team_id"] == "t-1"
assert payload["organization_id"] == "o-1"
assert payload["end_user"] == "tenant-a"
metadata = json.loads(payload["metadata"]) if isinstance(payload["metadata"], str) else payload["metadata"]
assert metadata["user_api_key"] == hash_token(RAW_KEY)
request_tags = (
json.loads(payload["request_tags"]) if isinstance(payload["request_tags"], str) else payload["request_tags"]
)
assert request_tags == ["batch:job-42"]
def test_idempotent_record_books_through_atomic_reserve_not_funnel():
deps = RecordingDeps(compute_cost_error=RuntimeError("pricing must not be consulted"))
result = run(process_external_usage_record(make_record(cost=0.123, idempotency_key="batch-1-line-1"), deps.as_deps()))
result = run(
process_external_usage_record(make_record(cost=0.123, idempotency_key="batch-1-line-1"), deps.as_deps())
)
assert result.status == "recorded"
assert result.spend == 0.123
assert result.request_id == "batch-1-line-1"
assert len(deps.spend_calls) == 1
assert len(deps.reserve_calls) == 1
reserve_call = deps.reserve_calls[0]
assert reserve_call["request_id"] == "batch-1-line-1"
assert reserve_call["hashed_token"] == hash_token(RAW_KEY)
assert reserve_call["cost"] == 0.123
assert reserve_call["key"] == DEFAULT_KEY
assert deps.spend_calls == []
assert deps.compute_cost_calls == []
call = deps.spend_calls[0]
assert call["token"] == hash_token(RAW_KEY)
assert call["token"] != RAW_KEY
assert call["user_id"] == "u-1"
assert call["team_id"] == "t-1"
assert call["org_id"] == "o-1"
assert call["response_cost"] == 0.123
response = call["completion_response"]
assert response.id == "batch-1-line-1"
assert response.usage.prompt_tokens == 100
assert response.usage.completion_tokens == 50
assert response.usage.total_tokens == 150
kwargs = call["kwargs"]
assert kwargs["call_type"] == "ingest_external_usage"
metadata = kwargs["litellm_params"]["metadata"]
assert metadata["user_api_key"] == hash_token(RAW_KEY)
assert metadata["user_api_key"] != RAW_KEY
def test_computed_cost_used_when_no_explicit_cost():
def test_computed_cost_is_resolved_before_reserving():
deps = RecordingDeps(compute_cost_result=0.07)
result = run(process_external_usage_record(make_record(idempotency_key="k-2"), deps.as_deps()))
assert result.status == "recorded"
assert result.spend == 0.07
assert len(deps.compute_cost_calls) == 1
assert deps.compute_cost_calls[0][1] == "gpt-4o-mini"
assert deps.spend_calls[0]["response_cost"] == 0.07
assert deps.reserve_calls[0]["cost"] == 0.07
def test_unpriceable_model_without_explicit_cost_is_error_not_zero_spend():
def test_unpriceable_model_without_explicit_cost_is_error_and_books_nothing():
deps = RecordingDeps(compute_cost_error=ValueError("unknown model"))
result = run(process_external_usage_record(make_record(idempotency_key="k-3"), deps.as_deps()))
assert result.status == "error"
assert result.spend is None
assert "explicit cost" in (result.error or "")
assert len(deps.spend_calls) == 0
assert deps.reserve_calls == []
assert deps.spend_calls == []
def test_unknown_key_is_rejected_and_never_books_spend():
def test_unknown_key_is_rejected_and_books_nothing():
deps = RecordingDeps(key=None)
result = run(process_external_usage_record(make_record(idempotency_key="k-4"), deps.as_deps()))
assert result.status == "error"
assert result.error == "api key not found"
assert len(deps.spend_calls) == 0
assert deps.reserve_calls == []
assert deps.spend_calls == []
def test_duplicate_idempotency_key_is_skipped_and_never_rebooks():
@ -143,16 +180,30 @@ def test_duplicate_idempotency_key_is_skipped_and_never_rebooks():
result = run(process_external_usage_record(make_record(idempotency_key="k-5"), deps.as_deps()))
assert result.status == "duplicate"
assert result.request_id == "k-5"
assert len(deps.spend_calls) == 0
assert len(deps.reserve_calls) == 1
assert deps.spend_calls == []
def test_missing_idempotency_key_generates_request_id_and_skips_dedup_probe():
def test_missing_idempotency_key_uses_funnel_with_generated_request_id():
deps = RecordingDeps()
result = run(process_external_usage_record(make_record(cost=0.01), deps.as_deps()))
assert result.status == "recorded"
assert result.request_id == GENERATED_ID
assert deps.spend_calls[0]["completion_response"].id == GENERATED_ID
assert deps.exists_calls == []
assert deps.reserve_calls == []
assert len(deps.spend_calls) == 1
call = deps.spend_calls[0]
assert call["token"] == hash_token(RAW_KEY)
assert call["user_id"] == "u-1"
assert call["team_id"] == "t-1"
assert call["org_id"] == "o-1"
assert call["response_cost"] == 0.01
response = call["completion_response"]
assert response.id == GENERATED_ID
assert response.usage.total_tokens == 150
metadata = call["kwargs"]["litellm_params"]["metadata"]
assert metadata["user_api_key"] == hash_token(RAW_KEY)
assert metadata["user_api_key"] != RAW_KEY
def test_end_time_defaults_to_start_time_and_end_before_start_rejected():
@ -167,17 +218,15 @@ def test_end_time_defaults_to_start_time_and_end_before_start_rejected():
make_record(start_time=start, end_time=earlier)
def test_tags_and_end_user_flow_into_payload():
def test_record_without_idempotency_key_still_flows_tags_to_funnel_kwargs():
deps = RecordingDeps()
result = run(
process_external_usage_record(
make_record(cost=0.01, idempotency_key="k-8", tags=["batch:job-42"], end_user_id="tenant-a"),
deps.as_deps(),
make_record(cost=0.01, tags=["batch:job-42"], end_user_id="tenant-a"), deps.as_deps()
)
)
assert result.status == "recorded"
call = deps.spend_calls[0]
metadata = call["kwargs"]["litellm_params"]["metadata"]
metadata = deps.spend_calls[0]["kwargs"]["litellm_params"]["metadata"]
assert metadata["tags"] == ["batch:job-42"]
assert metadata["user_api_key_end_user_id"] == "tenant-a"
assert call["end_user_id"] == "tenant-a"
assert deps.spend_calls[0]["end_user_id"] == "tenant-a"

View file

@ -12859,6 +12859,36 @@ export interface paths {
patch?: never;
trace?: never;
};
"/spend/usage": {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
get?: never;
put?: never;
/**
* Ingest External Usage
* @description PROXY_ADMIN ONLY: record externally measured usage into the same spend pipeline as proxy-routed traffic.
*
* For inference traffic that legitimately bypasses the proxy (for example async batch processors
* dispatching directly to model gateways), so budgets and spend stay coherent in litellm as the
* single metering system.
*
* Attribution (user/team/org) is derived from the given virtual key. Records accept an optional
* idempotency_key, stored as the spend-log request_id: the reservation insert, counter updates and
* dedup are checked atomically at the database primary key, so overlapping retries are safe. When
* cost is omitted it is computed from litellm pricing; records whose model cannot be priced are
* rejected with an error instead of being booked as zero spend.
*/
post: operations["ingest_external_usage_spend_usage_post"];
delete?: never;
options?: never;
head?: never;
patch?: never;
trace?: never;
};
"/spend/users": {
parameters: {
query?: never;
@ -22482,6 +22512,7 @@ export interface components {
/** ChatCompletionAudioObject */
ChatCompletionAudioObject: {
input_audio: components["schemas"]["InputAudio"];
prompt_cache_breakpoint?: components["schemas"]["PromptCacheBreakpoint"];
/**
* Type
* @constant
@ -24401,6 +24432,41 @@ export interface components {
/** Updated At */
updated_at?: number | null;
};
/** ExternalUsageRecord */
ExternalUsageRecord: {
/**
* Api Key
* @description Raw virtual key (sk-...) to attribute usage to. Never logged.
*/
api_key: string;
/** Completion Tokens */
completion_tokens: number;
/**
* Cost
* @description Explicit cost in USD. Computed from litellm pricing when omitted.
*/
cost?: number | null;
/** End Time */
end_time?: string | null;
/** End User Id */
end_user_id?: string | null;
/**
* Idempotency Key
* @description Becomes the spend-log request_id for dedup on retries.
*/
idempotency_key?: string | null;
/** Model */
model: string;
/** Prompt Tokens */
prompt_tokens: number;
/**
* Start Time
* Format: date-time
*/
start_time: string;
/** Tags */
tags?: string[] | null;
};
/**
* FacetListResponse
* @description The distinct values one column takes over a filtered query. `data` holds bare values, not entity rows.
@ -30582,6 +30648,19 @@ export interface components {
prompt_id: string;
prompt_info?: components["schemas"]["PromptInfo"] | null;
};
/**
* PromptCacheBreakpoint
* @description Marks the exact end of a reusable prompt prefix.
*
* The breakpoint inherits its TTL from the request's `prompt_cache_options.ttl`; the boundary is not rounded to a token block.
*/
PromptCacheBreakpoint: {
/**
* Mode
* @constant
*/
mode: "explicit";
};
/** PromptInfo */
PromptInfo: {
/**
@ -34118,6 +34197,30 @@ export interface components {
/** Type */
type: string;
};
/** UsageIngestRecordResult */
UsageIngestRecordResult: {
/** Error */
error?: string | null;
/** Request Id */
request_id: string;
/** Spend */
spend?: number | null;
/**
* Status
* @enum {string}
*/
status: "recorded" | "duplicate" | "error";
};
/** UsageIngestRequest */
UsageIngestRequest: {
/** Records */
records: components["schemas"]["ExternalUsageRecord"][];
};
/** UsageIngestResponse */
UsageIngestResponse: {
/** Results */
results: components["schemas"]["UsageIngestRecordResult"][];
};
/** UsageLogEntry */
UsageLogEntry: {
/** Action */
@ -50983,6 +51086,39 @@ export interface operations {
};
};
};
ingest_external_usage_spend_usage_post: {
parameters: {
query?: never;
header?: never;
path?: never;
cookie?: never;
};
requestBody: {
content: {
"application/json": components["schemas"]["UsageIngestRequest"];
};
};
responses: {
/** @description Successful Response */
200: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["UsageIngestResponse"];
};
};
/** @description Validation Error */
422: {
headers: {
[name: string]: unknown;
};
content: {
"application/json": components["schemas"]["HTTPValidationError"];
};
};
};
};
spend_user_fn_spend_users_get: {
parameters: {
query?: {