mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(proxy): add POST /spend/usage to ingest externally measured usage into spend tracking
This commit is contained in:
parent
732bba00df
commit
bdfd0e4691
4 changed files with 392 additions and 0 deletions
|
|
@ -632,6 +632,7 @@ class LiteLLMRoutes(enum.Enum):
|
||||||
"/jwt/key/mapping/delete",
|
"/jwt/key/mapping/delete",
|
||||||
"/jwt/key/mapping/list",
|
"/jwt/key/mapping/list",
|
||||||
"/jwt/key/mapping/info",
|
"/jwt/key/mapping/info",
|
||||||
|
"/spend/usage",
|
||||||
]
|
]
|
||||||
+ key_management_routes
|
+ key_management_routes
|
||||||
+ mcp_management_routes
|
+ mcp_management_routes
|
||||||
|
|
|
||||||
|
|
@ -549,6 +549,9 @@ from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||||
router as spend_management_router,
|
router as spend_management_router,
|
||||||
)
|
)
|
||||||
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
|
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.types_utils.utils import get_instance_fn
|
||||||
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
|
from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import (
|
||||||
router as ui_crud_endpoints_router,
|
router as ui_crud_endpoints_router,
|
||||||
|
|
@ -16449,6 +16452,7 @@ app.include_router(organization_router)
|
||||||
app.include_router(customer_router)
|
app.include_router(customer_router)
|
||||||
app.include_router(management_v1_router)
|
app.include_router(management_v1_router)
|
||||||
app.include_router(spend_management_router)
|
app.include_router(spend_management_router)
|
||||||
|
app.include_router(usage_ingestion_router)
|
||||||
app.include_router(caching_router)
|
app.include_router(caching_router)
|
||||||
app.include_router(analytics_router)
|
app.include_router(analytics_router)
|
||||||
app.include_router(callback_management_endpoints_router)
|
app.include_router(callback_management_endpoints_router)
|
||||||
|
|
|
||||||
204
litellm/proxy/spend_tracking/usage_ingestion_endpoints.py
Normal file
204
litellm/proxy/spend_tracking/usage_ingestion_endpoints.py
Normal 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)
|
||||||
|
|
@ -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"
|
||||||
Loading…
Add table
Reference in a new issue