diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7d6829aca70..dfe6bc672bf 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e343d46f872..56169bed74c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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) diff --git a/litellm/proxy/spend_tracking/usage_ingestion_endpoints.py b/litellm/proxy/spend_tracking/usage_ingestion_endpoints.py new file mode 100644 index 00000000000..8e3e3dd448d --- /dev/null +++ b/litellm/proxy/spend_tracking/usage_ingestion_endpoints.py @@ -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) diff --git a/tests/test_litellm/proxy/spend_tracking/test_usage_ingestion_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_usage_ingestion_endpoints.py new file mode 100644 index 00000000000..360f2e1af6e --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_usage_ingestion_endpoints.py @@ -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"