feat(proxy): add POST /spend/usage to ingest externally measured usage into spend tracking

This commit is contained in:
todayim 2026-08-05 16:50:25 +08:00
parent 732bba00df
commit bdfd0e4691
4 changed files with 392 additions and 0 deletions

View file

@ -632,6 +632,7 @@ class LiteLLMRoutes(enum.Enum):
"/jwt/key/mapping/delete",
"/jwt/key/mapping/list",
"/jwt/key/mapping/info",
"/spend/usage",
]
+ key_management_routes
+ mcp_management_routes

View file

@ -549,6 +549,9 @@ from litellm.proxy.spend_tracking.spend_management_endpoints import (
router as spend_management_router,
)
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
from litellm.proxy.spend_tracking.usage_ingestion_endpoints import (
router as usage_ingestion_router,
)
from litellm.proxy.types_utils.utils import get_instance_fn
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
router as ui_crud_endpoints_router,
@ -16449,6 +16452,7 @@ app.include_router(organization_router)
app.include_router(customer_router)
app.include_router(management_v1_router)
app.include_router(spend_management_router)
app.include_router(usage_ingestion_router)
app.include_router(caching_router)
app.include_router(analytics_router)
app.include_router(callback_management_endpoints_router)

View file

@ -0,0 +1,204 @@
import uuid
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import datetime
from typing import Final, Literal, NamedTuple
from fastapi import APIRouter, Depends, 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.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.utils import hash_token
router: Final = APIRouter()
MAX_RECORDS_PER_REQUEST: Final = 1000
class ExternalUsageRecord(BaseModel):
api_key: str = Field(min_length=1, description="Raw virtual key (sk-...) to attribute usage to. Never logged.")
model: str = Field(min_length=1)
prompt_tokens: int = Field(ge=0)
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.")
tags: list[str] | None = None
end_user_id: str | None = None
@model_validator(mode="after")
def end_time_not_before_start_time(self) -> "ExternalUsageRecord":
if self.end_time is not None and self.end_time < self.start_time:
raise ValueError("end_time must not be before start_time")
return self
class UsageIngestRequest(BaseModel):
records: list[ExternalUsageRecord] = Field(min_length=1, max_length=MAX_RECORDS_PER_REQUEST)
class UsageIngestRecordResult(BaseModel):
request_id: str
status: Literal["recorded", "duplicate", "error"]
spend: float | None = None
error: str | None = None
class UsageIngestResponse(BaseModel):
results: list[UsageIngestRecordResult]
class KeyAttribution(NamedTuple):
user_id: str | None
team_id: str | None
organization_id: str | None
RecordSpendFn = Callable[..., Awaitable[None]]
@dataclass(frozen=True, slots=True)
class UsageIngestionDeps:
lookup_key: Callable[[str], Awaitable[KeyAttribution | None]]
spend_log_exists: Callable[[str], Awaitable[bool]]
record_spend: RecordSpendFn
compute_cost: Callable[[litellm.ModelResponse, str], float]
generate_request_id: Callable[[], str]
def _build_usage_kwargs(record: ExternalUsageRecord, hashed_token: str) -> dict[str, object]:
metadata = {
"user_api_key": hashed_token,
"user_api_key_end_user_id": record.end_user_id,
"tags": list(record.tags) if record.tags else [],
}
return {
"model": record.model,
"call_type": "ingest_external_usage",
"litellm_params": {"model": record.model, "metadata": metadata},
}
def _build_completion_response(record: ExternalUsageRecord, request_id: str) -> litellm.ModelResponse:
total_tokens: Final = record.prompt_tokens + record.completion_tokens
usage: Final = litellm.Usage(
prompt_tokens=record.prompt_tokens,
completion_tokens=record.completion_tokens,
total_tokens=total_tokens,
)
return litellm.ModelResponse(
id=request_id,
model=record.model,
created=int(record.start_time.timestamp()),
usage=usage,
)
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:
request_id: Final = record.idempotency_key or deps.generate_request_id()
hashed_token: Final = hash_token(record.api_key)
key: Final = await deps.lookup_key(hashed_token)
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
verbose_proxy_logger.info("ingest usage: cost computation failed for model %s: %s", record.model, e)
return UsageIngestRecordResult(
request_id=request_id,
status="error",
error="could not compute cost for this model, pass an explicit cost",
)
await deps.record_spend(
token=hashed_token,
user_id=key.user_id,
end_user_id=record.end_user_id,
team_id=key.team_id,
kwargs=_build_usage_kwargs(record, hashed_token),
completion_response=response,
start_time=record.start_time,
end_time=record.end_time or record.start_time,
response_cost=cost,
org_id=key.organization_id,
)
return UsageIngestRecordResult(request_id=request_id, status="recorded", spend=cost)
def _attribution_of(key_row: object) -> KeyAttribution:
return KeyAttribution(
user_id=getattr(key_row, "user_id", None),
team_id=getattr(key_row, "team_id", None),
organization_id=getattr(key_row, "organization_id", None),
)
def default_ingestion_deps() -> UsageIngestionDeps:
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj
if prisma_client is None:
raise ProxyException(
message="Prisma Client is not initialized",
type="internal_error",
param="None",
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
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)
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
return UsageIngestionDeps(
lookup_key=lookup_key,
spend_log_exists=spend_log_exists,
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()),
)
@router.post(
"/spend/usage",
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth)],
response_model=UsageIngestResponse,
)
async def ingest_external_usage(request: UsageIngestRequest) -> UsageIngestResponse:
"""
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, 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.
"""
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

@ -0,0 +1,183 @@
import asyncio
import os
import sys
from datetime import datetime, timezone
from typing import Any
import pytest
from pydantic import ValidationError
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.proxy.spend_tracking.usage_ingestion_endpoints import (
ExternalUsageRecord,
KeyAttribution,
UsageIngestionDeps,
process_external_usage_record,
)
from litellm.proxy.utils import hash_token
RAW_KEY = "sk-test-batch-dispatch-key"
GENERATED_ID = "generated-uuid-1"
DEFAULT_KEY = KeyAttribution(user_id="u-1", team_id="t-1", organization_id="o-1")
class RecordingDeps:
def __init__(
self,
key: KeyAttribution | None = DEFAULT_KEY,
existing_ids: frozenset[str] = frozenset(),
compute_cost_result: float = 0.05,
compute_cost_error: 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.spend_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 record_spend(**kwargs: Any) -> None:
self.spend_calls.append(kwargs)
def compute_cost(response: Any, model: str) -> float:
self.compute_cost_calls.append((response, model))
if self._compute_cost_error is not None:
raise self._compute_cost_error
return self._compute_cost_result
return UsageIngestionDeps(
lookup_key=lookup_key,
spend_log_exists=spend_log_exists,
record_spend=record_spend,
compute_cost=compute_cost,
generate_request_id=lambda: GENERATED_ID,
)
def make_record(**overrides: Any) -> ExternalUsageRecord:
base: dict[str, Any] = {
"api_key": RAW_KEY,
"model": "gpt-4o-mini",
"prompt_tokens": 100,
"completion_tokens": 50,
"start_time": datetime(2026, 8, 5, 12, 0, 0, tzinfo=timezone.utc),
}
base.update(overrides)
return ExternalUsageRecord(**base)
def run(coro: Any) -> Any:
return asyncio.run(coro)
def test_records_spend_with_explicit_cost_without_calling_pricing():
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()))
assert result.status == "recorded"
assert result.spend == 0.123
assert result.request_id == "batch-1-line-1"
assert len(deps.spend_calls) == 1
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():
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
def test_unpriceable_model_without_explicit_cost_is_error_not_zero_spend():
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
def test_unknown_key_is_rejected_and_never_books_spend():
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
def test_duplicate_idempotency_key_is_skipped_and_never_rebooks():
deps = RecordingDeps(existing_ids=frozenset({"k-5"}))
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
def test_missing_idempotency_key_generates_request_id_and_skips_dedup_probe():
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 == []
def test_end_time_defaults_to_start_time_and_end_before_start_rejected():
deps = RecordingDeps()
start = datetime(2026, 8, 5, 12, 0, 0, tzinfo=timezone.utc)
run(process_external_usage_record(make_record(cost=0.01, start_time=start), deps.as_deps()))
assert deps.spend_calls[0]["start_time"] == start
assert deps.spend_calls[0]["end_time"] == start
earlier = datetime(2026, 8, 5, 11, 0, 0, tzinfo=timezone.utc)
with pytest.raises(ValidationError):
make_record(start_time=start, end_time=earlier)
def test_tags_and_end_user_flow_into_payload():
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(),
)
)
assert result.status == "recorded"
call = deps.spend_calls[0]
metadata = call["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"