From 838c7a7ea72b7ba98428b4734dcec7ceab9231d2 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 24 Jul 2026 18:10:15 -0700 Subject: [PATCH 1/3] feat(passthrough): record upstream-reported cost and token usage A pass-through target that fans a single HTTP request out to several models internally cannot be priced from its response body, so LiteLLM had nothing to record and every such request landed in the spend logs with zero cost and zero tokens. The target now reports the totals for the whole request in x-litellm-response-cost and x-litellm-total-tokens response headers, and LiteLLM records those values as-is rather than recomputing them. The headers are read on every upstream response, so a request that burned tokens before failing still books its spend on the failure row instead of being dropped for having a 4xx/5xx status. Only what the upstream actually reported is written, so a target that sends a cost but no token count keeps the token count LiteLLM derived on its own; a target that sends neither header is untouched, which is the normal case for Anthropic, Vertex and friends. Two supporting fixes fall out of this. The rate limiter only pulled token counts off response shapes it models, so pass-through usage never charged the TPM window and a team could exceed its shared token limit through pass-through traffic alone; it now falls back to combined_usage_object. And the streaming success path reset response_cost unconditionally before the assembled response recomputed it, which discarded any cost a pass-through handler had already established (the pass-through branch right below it has always intended to preserve exactly that). --- litellm/litellm_core_utils/litellm_logging.py | 9 +- .../hooks/parallel_request_limiter_v3.py | 14 +- .../pass_through_endpoints.py | 26 +++ .../upstream_usage_headers.py | 120 ++++++++++ .../hooks/test_parallel_request_limiter_v3.py | 50 ++++ .../test_pass_through_endpoints.py | 214 +++++++++++++++++- .../test_upstream_usage_headers.py | 131 +++++++++++ 7 files changed, 559 insertions(+), 5 deletions(-) create mode 100644 litellm/proxy/pass_through_endpoints/upstream_usage_headers.py create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/test_upstream_usage_headers.py diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index c9e70b7db73..e1e5499d008 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1878,7 +1878,14 @@ class Logging(LiteLLMLoggingBaseClass): elif standard_logging_object is not None: self.model_call_details["standard_logging_object"] = standard_logging_object else: - self.model_call_details["response_cost"] = None + # Streaming reaches here before its cost is known, so the cost + # is seeded to None, but only when nothing has already + # established one. A stream that assembles into a response + # object recomputes the cost right after this; a pass-through + # stream cannot (its body is opaque) and carries the cost its + # upstream reported in the response headers, which an + # unconditional reset would discard. + self.model_call_details.setdefault("response_cost", None) result = self._transform_usage_objects(result=result) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b2216488db2..81d46247507 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -3086,8 +3086,12 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): model_group = get_model_group_from_litellm_kwargs(kwargs) - # Get total tokens from response - total_tokens = 0 + # Get total tokens from response. Responses LiteLLM does not model + # (e.g. pass-through, whose usage is reported by the upstream rather + # than parsed out of the body) carry their usage in + # ``combined_usage_object`` instead, and would otherwise never charge + # the TPM window. + _usage: Union[Usage, dict, None] = None if isinstance( response_obj, ( @@ -3098,7 +3102,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ), ): _usage = getattr(response_obj, "usage", None) - total_tokens = self._get_total_tokens_from_usage(usage=_usage, rate_limit_type=rate_limit_type) + else: + _combined_usage = kwargs.get("combined_usage_object") + if isinstance(_combined_usage, Usage): + _usage = _combined_usage + total_tokens = self._get_total_tokens_from_usage(usage=_usage, rate_limit_type=rate_limit_type) reserved_tokens = self._get_reserved_tokens_from_kwargs( kwargs=kwargs, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 9364d7eae3a..9c5c545fe76 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -77,9 +77,14 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( EndpointType, PassthroughStandardLoggingPayload, ) +from litellm.types.utils import Usage from .streaming_handler import PassThroughStreamingHandler from .success_handler import PassThroughEndpointLogging +from .upstream_usage_headers import ( + UpstreamReportedUsage, + apply_upstream_reported_usage, +) router = APIRouter() @@ -687,12 +692,17 @@ def _build_passthrough_failure_request_payload( kwargs: Optional[dict], logging_obj: Optional[LiteLLMLoggingObj], custom_llm_provider: Optional[str], + upstream_usage: UpstreamReportedUsage | None = None, ) -> dict: """Build the ``request_data`` dict passed to ``post_call_failure_hook``. Shared by the outer exception handler (LiteLLM-internal failures) and upstream HTTP error logging, so both failure paths report the same shape of request data (model, custom_llm_provider, litellm_logging_obj, ...). + + ``upstream_usage`` carries the cost and tokens an upstream reported on an + error response. Spend tracking only attributes a recovered cost when it + comes paired with a usage object, so both keys are written together. """ request_payload: dict = dict(parsed_body or {}) if kwargs: @@ -703,6 +713,9 @@ def _build_passthrough_failure_request_payload( request_payload["model"] = parsed_body.get("model", "") if "custom_llm_provider" not in request_payload and custom_llm_provider: request_payload["custom_llm_provider"] = custom_llm_provider + if upstream_usage is not None: + request_payload["response_cost"] = upstream_usage.response_cost or 0.0 + request_payload["combined_usage_object"] = Usage(total_tokens=upstream_usage.total_tokens or 0) return request_payload @@ -1125,6 +1138,11 @@ async def pass_through_request( response = await async_client.send(req, stream=stream) + upstream_usage = apply_upstream_reported_usage( + logging_obj=logging_obj, + headers=response.headers, + ) + await _log_passthrough_upstream_failure( response=response, user_api_key_dict=user_api_key_dict, @@ -1133,6 +1151,7 @@ async def pass_through_request( kwargs=kwargs, logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, + upstream_usage=upstream_usage, ), ) @@ -1187,6 +1206,11 @@ async def pass_through_request( ) verbose_proxy_logger.debug("response.headers= %s", response.headers) + upstream_usage = apply_upstream_reported_usage( + logging_obj=logging_obj, + headers=response.headers, + ) + if _is_streaming_response(response) is True: logging_obj.stream = True logging_obj.model_call_details["stream"] = True @@ -1199,6 +1223,7 @@ async def pass_through_request( kwargs=kwargs, logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, + upstream_usage=upstream_usage, ), ) @@ -1278,6 +1303,7 @@ async def pass_through_request( kwargs=kwargs, logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, + upstream_usage=upstream_usage, ) failure_request_payload["response_body"] = response_body await _log_passthrough_upstream_failure( diff --git a/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py new file mode 100644 index 00000000000..6531cff1c1d --- /dev/null +++ b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py @@ -0,0 +1,120 @@ +"""Upstream-reported cost and usage for pass-through endpoints. + +Some pass-through targets invoke several models internally, so LiteLLM cannot +price the request from the response body. Those targets report the totals for +the whole HTTP request in ``x-litellm-*`` response headers instead; LiteLLM +records the reported values without recomputing them. +""" + +import math +from dataclasses import dataclass + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.types.utils import Usage + +UPSTREAM_RESPONSE_COST_HEADER = "x-litellm-response-cost" +UPSTREAM_TOTAL_TOKENS_HEADER = "x-litellm-total-tokens" + + +@dataclass(frozen=True, slots=True) +class UpstreamReportedUsage: + """Totals an upstream pass-through target reported for one HTTP request. + + ``None`` means the upstream did not report a usable value, either because + the header was absent or because it could not be parsed. + """ + + response_cost: float | None + total_tokens: int | None + + +def _parse_response_cost(raw_value: str | None) -> float | None: + if raw_value is None: + verbose_proxy_logger.warning( + "pass_through_endpoint: upstream did not send %s; recording 0 cost for this request", + UPSTREAM_RESPONSE_COST_HEADER, + ) + return None + try: + response_cost = float(raw_value) + except ValueError: + verbose_proxy_logger.warning( + "pass_through_endpoint: upstream sent unparseable %s=%r; recording 0 cost for this request", + UPSTREAM_RESPONSE_COST_HEADER, + raw_value, + ) + return None + if not math.isfinite(response_cost) or response_cost < 0: + verbose_proxy_logger.warning( + "pass_through_endpoint: upstream sent out-of-range %s=%r; recording 0 cost for this request", + UPSTREAM_RESPONSE_COST_HEADER, + raw_value, + ) + return None + return response_cost + + +def _parse_total_tokens(raw_value: str | None) -> int | None: + if raw_value is None: + verbose_proxy_logger.warning( + "pass_through_endpoint: upstream did not send %s; recording 0 tokens for this request", + UPSTREAM_TOTAL_TOKENS_HEADER, + ) + return None + try: + total_tokens = int(raw_value) + except ValueError: + verbose_proxy_logger.warning( + "pass_through_endpoint: upstream sent unparseable %s=%r; recording 0 tokens for this request", + UPSTREAM_TOTAL_TOKENS_HEADER, + raw_value, + ) + return None + if total_tokens < 0: + verbose_proxy_logger.warning( + "pass_through_endpoint: upstream sent negative %s=%r; recording 0 tokens for this request", + UPSTREAM_TOTAL_TOKENS_HEADER, + raw_value, + ) + return None + return total_tokens + + +def parse_upstream_reported_usage(headers: httpx.Headers) -> UpstreamReportedUsage | None: + """Read the reported totals off an upstream pass-through response. + + Returns ``None`` when neither header is present, which is the normal case + for a target that does not speak this contract (e.g. Anthropic or Vertex, + whose cost LiteLLM derives from the response body instead). + """ + raw_response_cost = headers.get(UPSTREAM_RESPONSE_COST_HEADER) + raw_total_tokens = headers.get(UPSTREAM_TOTAL_TOKENS_HEADER) + if raw_response_cost is None and raw_total_tokens is None: + return None + return UpstreamReportedUsage( + response_cost=_parse_response_cost(raw_response_cost), + total_tokens=_parse_total_tokens(raw_total_tokens), + ) + + +def apply_upstream_reported_usage( + logging_obj: LiteLLMLoggingObj, + headers: httpx.Headers, +) -> UpstreamReportedUsage | None: + """Record the upstream's reported totals on the request's logging object. + + Only the values the upstream actually reported are written, so a target + that reports cost but not tokens keeps the token count LiteLLM derived on + its own rather than having it zeroed. + """ + reported = parse_upstream_reported_usage(headers) + if reported is None: + return None + if reported.response_cost is not None: + logging_obj.model_call_details["response_cost"] = reported.response_cost + if reported.total_tokens is not None: + logging_obj.model_call_details["combined_usage_object"] = Usage(total_tokens=reported.total_tokens) + return reported diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index c76e1a60afd..27b1c025fbb 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -4652,3 +4652,53 @@ async def test_streaming_mirror_matches_non_streaming_header_shape(monkeypatch): f" non_streaming={_rl_only(non_stream_headers)}" ) assert "x-ratelimit-model_per_key-remaining-requests" in stream_slp_headers + + +@pytest.mark.asyncio +async def test_async_log_success_event_counts_passthrough_reported_tokens(monkeypatch): + """ + Pass-through responses are not modelled as ModelResponse/EmbeddingResponse, + so their usage rides on ``combined_usage_object``. Without honoring it, a + pass-through request never charged the TPM window and a team could exceed + its shared token limit purely through pass-through traffic. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + + _api_key = hash_token("sk-passthrough") + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + monkeypatch.setattr( + parallel_request_handler, "get_rate_limit_type", lambda: "total" + ) + + captured_operations = [] + + async def mock_increment_pipeline(increment_list, **kwargs): + captured_operations.extend(increment_list) + return True + + monkeypatch.setattr( + parallel_request_handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + mock_increment_pipeline, + ) + + await parallel_request_handler.async_log_success_event( + kwargs={ + "standard_logging_object": { + "metadata": {"user_api_key_hash": _api_key, "user_api_key_team_id": "team-fil"} + }, + "combined_usage_object": Usage(total_tokens=1874), + }, + response_obj={"response": "upstream body was never parsed"}, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + token_operations = [op for op in captured_operations if op["key"].endswith(":tokens")] + assert token_operations, "pass-through usage should charge the TPM counter" + assert all(op["increment_value"] == 1874 for op in token_operations) + assert any("team-fil" in op["key"] for op in token_operations), ( + "team TPM window must be charged so pass-through shares the team's limit" + ) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index fdf629c36bd..10706b4ea64 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -3,7 +3,7 @@ import json import logging import os import sys -from contextlib import ExitStack +from contextlib import ExitStack, contextmanager from io import BytesIO from types import SimpleNamespace from typing import Optional @@ -31,6 +31,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( resolve_pass_through_request_timeout, resolve_llm_passthrough_timeout, ) +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -4617,3 +4618,214 @@ async def test_pass_through_relay_full_consumption_logs_no_partial_relay_warning finally: cleanup() await fake_client.aclose() + + +class _StandardLoggingPayloadRecorder(CustomLogger): + """Stands in for a real logging integration so the assertions run against + the StandardLoggingPayload that spend tracking consumes, not an internal.""" + + def __init__(self): + super().__init__() + self.payloads = [] + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + self.payloads.append(kwargs.get("standard_logging_object")) + + +@contextmanager +def _recording_success_callback(): + import litellm + + recorder = _StandardLoggingPayloadRecorder() + original = litellm._async_success_callback + litellm._async_success_callback = [*original, recorder] + try: + yield recorder + finally: + litellm._async_success_callback = original + + +def _enter_upstream_usage_mocks(stack, parsed_body): + """Same seams as _enter_relay_logging_mocks, but leaves the real + pass-through success handler in place and captures the coroutines the + logging worker would have run so the test can await them.""" + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + mock_proxy_logging = stack.enter_context( + patch("litellm.proxy.proxy_server.proxy_logging_obj") + ) + mock_proxy_logging.pre_call_hook = AsyncMock(return_value=parsed_body) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + mock_proxy_logging.get_proxy_hook = MagicMock(return_value=None) + + enqueued = [] + stack.enter_context( + patch.object( + GLOBAL_LOGGING_WORKER, + "ensure_initialized_and_enqueue", + new=MagicMock(side_effect=lambda async_coroutine: enqueued.append(async_coroutine)), + ) + ) + return mock_proxy_logging, enqueued + + +async def _run_upstream_reporting_passthrough(upstream_headers, status_code=200): + """Drive a generic pass-through against an upstream that reports its own + cost/usage. Returns (recorded standard logging payloads, proxy logging mock).""" + from litellm.proxy._types import UserAPIKeyAuth + + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=status_code, + headers={"content-type": "application/json", **upstream_headers}, + stream=_RecordingUpstreamByteStream((b'{"answer": "ok"}',)), + ), + timeout=None, + ) + try: + with ExitStack() as stack: + mock_proxy_logging, enqueued = _enter_upstream_usage_mocks(stack, {}) + with _recording_success_callback() as recorder: + await pass_through_request( + request=_relay_client_request(method="POST"), + target="http://internal-api.test/v1/summarize", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth( + api_key="sk-upstream-usage", team_id="team-fil" + ), + ) + for coroutine in enqueued: + await coroutine + return recorder.payloads, mock_proxy_logging + finally: + cleanup() + await fake_client.aclose() + + +@pytest.mark.asyncio +async def test_passthrough_records_cost_and_tokens_reported_by_upstream(): + """ + A pass-through target that prices the request itself reports the totals in + x-litellm-response-cost / x-litellm-total-tokens. LiteLLM must record those + values verbatim; before this, a generic pass-through always logged 0 spend + and 0 tokens because there was nothing in the body to price. + """ + payloads, _ = await _run_upstream_reporting_passthrough( + { + "x-litellm-response-cost": "0.000415", + "x-litellm-total-tokens": "1874", + } + ) + + assert len(payloads) == 1 + assert payloads[0]["response_cost"] == 0.000415 + assert payloads[0]["total_tokens"] == 1874 + + +@pytest.mark.asyncio +async def test_passthrough_records_zero_when_upstream_reports_zero(): + payloads, _ = await _run_upstream_reporting_passthrough( + {"x-litellm-response-cost": "0", "x-litellm-total-tokens": "0"} + ) + + assert len(payloads) == 1 + assert payloads[0]["response_cost"] == 0.0 + assert payloads[0]["total_tokens"] == 0 + + +@pytest.mark.asyncio +async def test_passthrough_records_zero_when_upstream_reports_nothing(): + payloads, _ = await _run_upstream_reporting_passthrough({}) + + assert len(payloads) == 1 + assert payloads[0]["response_cost"] == 0.0 + assert payloads[0]["total_tokens"] == 0 + + +@pytest.mark.asyncio +async def test_passthrough_records_upstream_reported_cost_on_error_response(): + """ + An upstream that already burned tokens before failing still reports them on + the error response, so the spend must land on the failure row rather than + being dropped because the status code was >= 400. + """ + import litellm + + _, mock_proxy_logging = await _run_upstream_reporting_passthrough( + { + "x-litellm-response-cost": "0.00021", + "x-litellm-total-tokens": "930", + }, + status_code=500, + ) + + mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + request_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs[ + "request_data" + ] + assert request_data["response_cost"] == 0.00021 + assert request_data["combined_usage_object"] == litellm.Usage(total_tokens=930) + + +@pytest.mark.asyncio +async def test_passthrough_error_response_without_usage_headers_records_no_spend(): + _, mock_proxy_logging = await _run_upstream_reporting_passthrough( + {}, status_code=500 + ) + + mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + request_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs[ + "request_data" + ] + assert "combined_usage_object" not in request_data + + +@pytest.mark.asyncio +async def test_streaming_passthrough_records_cost_and_tokens_reported_by_upstream(): + """ + The totals are reported in the response headers, which are known before the + first byte of the stream, so a streamed pass-through must record the same + spend a buffered one does. + """ + from fastapi.responses import StreamingResponse + + from litellm.proxy._types import UserAPIKeyAuth + + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={ + "content-type": "text/event-stream", + "x-litellm-response-cost": "0.00312", + "x-litellm-total-tokens": "4021", + }, + stream=_RecordingUpstreamByteStream((b'data: {"delta": "hi"}\n\n',)), + ), + timeout=None, + ) + try: + with ExitStack() as stack: + _, enqueued = _enter_upstream_usage_mocks(stack, {}) + with _recording_success_callback() as recorder: + response = await pass_through_request( + request=_relay_client_request(method="POST"), + target="http://internal-api.test/v1/summarize", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth( + api_key="sk-upstream-usage", team_id="team-fil" + ), + ) + assert isinstance(response, StreamingResponse) + assert [chunk async for chunk in response.body_iterator] == [ + b'data: {"delta": "hi"}\n\n' + ] + for coroutine in enqueued: + await coroutine + finally: + cleanup() + await fake_client.aclose() + + assert len(recorder.payloads) == 1 + assert recorder.payloads[0]["response_cost"] == 0.00312 + assert recorder.payloads[0]["total_tokens"] == 4021 diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_upstream_usage_headers.py b/tests/test_litellm/proxy/pass_through_endpoints/test_upstream_usage_headers.py new file mode 100644 index 00000000000..34c345b620f --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_upstream_usage_headers.py @@ -0,0 +1,131 @@ +import os +import sys + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy.pass_through_endpoints.upstream_usage_headers import ( + UpstreamReportedUsage, + apply_upstream_reported_usage, + parse_upstream_reported_usage, +) +from litellm.types.utils import Usage + + +def _headers(**values: str) -> httpx.Headers: + return httpx.Headers(values) + + +def test_parse_reads_both_totals(): + reported = parse_upstream_reported_usage( + _headers( + **{ + "x-litellm-response-cost": "0.000415", + "x-litellm-total-tokens": "1874", + } + ) + ) + + assert reported == UpstreamReportedUsage(response_cost=0.000415, total_tokens=1874) + + +def test_parse_returns_none_when_upstream_does_not_speak_the_contract(): + assert parse_upstream_reported_usage(_headers(**{"content-type": "application/json"})) is None + + +def test_parse_accepts_explicit_zero_totals(): + reported = parse_upstream_reported_usage( + _headers(**{"x-litellm-response-cost": "0", "x-litellm-total-tokens": "0"}) + ) + + assert reported == UpstreamReportedUsage(response_cost=0.0, total_tokens=0) + + +@pytest.mark.parametrize( + "raw_cost", + ["not-a-number", "-0.5", "nan", "inf", ""], +) +def test_parse_rejects_unusable_cost_but_keeps_tokens(raw_cost: str): + reported = parse_upstream_reported_usage( + _headers(**{"x-litellm-response-cost": raw_cost, "x-litellm-total-tokens": "12"}) + ) + + assert reported == UpstreamReportedUsage(response_cost=None, total_tokens=12) + + +@pytest.mark.parametrize("raw_tokens", ["1.5", "twelve", "-3", ""]) +def test_parse_rejects_unusable_tokens_but_keeps_cost(raw_tokens: str): + reported = parse_upstream_reported_usage( + _headers(**{"x-litellm-response-cost": "1.25", "x-litellm-total-tokens": raw_tokens}) + ) + + assert reported == UpstreamReportedUsage(response_cost=1.25, total_tokens=None) + + +def test_parse_reports_missing_counterpart_header(): + assert parse_upstream_reported_usage(_headers(**{"x-litellm-response-cost": "2.5"})) == UpstreamReportedUsage( + response_cost=2.5, total_tokens=None + ) + assert parse_upstream_reported_usage(_headers(**{"x-litellm-total-tokens": "7"})) == UpstreamReportedUsage( + response_cost=None, total_tokens=7 + ) + + +def _logging_obj() -> LiteLLMLoggingObj: + logging_obj = LiteLLMLoggingObj( + model="unknown", + messages=[{"role": "user", "content": "x"}], + stream=False, + call_type="pass_through_endpoint", + start_time=None, + litellm_call_id="test-call-id", + function_id="1", + ) + return logging_obj + + +def test_apply_records_reported_totals(): + logging_obj = _logging_obj() + + reported = apply_upstream_reported_usage( + logging_obj=logging_obj, + headers=_headers( + **{ + "x-litellm-response-cost": "0.000415", + "x-litellm-total-tokens": "1874", + } + ), + ) + + assert reported is not None + assert logging_obj.model_call_details["response_cost"] == 0.000415 + assert logging_obj.model_call_details["combined_usage_object"] == Usage(total_tokens=1874) + + +def test_apply_leaves_litellm_derived_values_alone_when_upstream_is_silent(): + logging_obj = _logging_obj() + logging_obj.model_call_details["response_cost"] = 9.99 + logging_obj.model_call_details["combined_usage_object"] = Usage(total_tokens=42) + + assert apply_upstream_reported_usage(logging_obj=logging_obj, headers=_headers()) is None + assert logging_obj.model_call_details["response_cost"] == 9.99 + assert logging_obj.model_call_details["combined_usage_object"] == Usage(total_tokens=42) + + +def test_apply_only_overwrites_what_upstream_reported(): + """A target that reports cost but not tokens must not zero out the token + count LiteLLM derived from the response body itself.""" + logging_obj = _logging_obj() + logging_obj.model_call_details["response_cost"] = 9.99 + logging_obj.model_call_details["combined_usage_object"] = Usage(total_tokens=42) + + apply_upstream_reported_usage( + logging_obj=logging_obj, + headers=_headers(**{"x-litellm-response-cost": "0.5"}), + ) + + assert logging_obj.model_call_details["response_cost"] == 0.5 + assert logging_obj.model_call_details["combined_usage_object"] == Usage(total_tokens=42) From ab44c8a8ce76dd5a99df94851687fd6bc4547fa4 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 24 Jul 2026 18:26:12 -0700 Subject: [PATCH 2/3] fix(passthrough): let an upstream-reported cost outrank cost_per_request PassThroughGenericEndpoint.cost_per_request defaults to 0.0, so every config-defined endpoint forwards a flat 0.0 even when the operator never configured one, and the success handler applied it over whatever cost was already established. That silently zeroed the cost an upstream reported for the request. The flat value is an estimate for targets LiteLLM cannot price, so it now yields to a target that priced the request itself; it still applies unchanged when no cost was reported. --- .../pass_through_endpoints/success_handler.py | 9 +++++ .../upstream_usage_headers.py | 14 ++++++++ .../test_pass_through_endpoints.py | 33 ++++++++++++++++++- 3 files changed, 55 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 932216141fe..84a956f0c29 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -30,6 +30,7 @@ from .llm_provider_handlers.gemini_passthrough_logging_handler import ( from .llm_provider_handlers.vertex_passthrough_logging_handler import ( VertexPassthroughLoggingHandler, ) +from .upstream_usage_headers import upstream_reported_cost cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler() @@ -448,10 +449,18 @@ class PassThroughEndpointLogging: Only set the cost per request if it's set in the passthrough logging payload. If it's not set, don't set it in the logging object. + + An upstream that priced the request itself always wins: ``cost_per_request`` + is a flat per-request estimate for targets LiteLLM cannot price, and it + defaults to 0.0 on every config-defined endpoint, so honoring it here + would zero out the real cost the upstream reported. """ ######################################################### # Check if cost per request is set ######################################################### + if upstream_reported_cost(logging_obj) is not None: + return kwargs + if passthrough_logging_payload.get("cost_per_request") is not None: kwargs["response_cost"] = passthrough_logging_payload.get("cost_per_request") logging_obj.model_call_details["response_cost"] = passthrough_logging_payload.get("cost_per_request") diff --git a/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py index 6531cff1c1d..e147dcb9547 100644 --- a/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py +++ b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py @@ -18,6 +18,11 @@ from litellm.types.utils import Usage UPSTREAM_RESPONSE_COST_HEADER = "x-litellm-response-cost" UPSTREAM_TOTAL_TOKENS_HEADER = "x-litellm-total-tokens" +# model_call_details key holding what the upstream reported, so later stages of +# the success path can tell an upstream-reported cost apart from one LiteLLM +# derived itself. +UPSTREAM_REPORTED_USAGE_KEY = "_litellm_upstream_reported_usage" + @dataclass(frozen=True, slots=True) class UpstreamReportedUsage: @@ -113,8 +118,17 @@ def apply_upstream_reported_usage( reported = parse_upstream_reported_usage(headers) if reported is None: return None + logging_obj.model_call_details[UPSTREAM_REPORTED_USAGE_KEY] = reported if reported.response_cost is not None: logging_obj.model_call_details["response_cost"] = reported.response_cost if reported.total_tokens is not None: logging_obj.model_call_details["combined_usage_object"] = Usage(total_tokens=reported.total_tokens) return reported + + +def upstream_reported_cost(logging_obj: LiteLLMLoggingObj) -> float | None: + """The cost the upstream reported for this request, if it reported one.""" + reported = logging_obj.model_call_details.get(UPSTREAM_REPORTED_USAGE_KEY) + if not isinstance(reported, UpstreamReportedUsage): + return None + return reported.response_cost diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 10706b4ea64..60340cfa47b 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -4670,7 +4670,9 @@ def _enter_upstream_usage_mocks(stack, parsed_body): return mock_proxy_logging, enqueued -async def _run_upstream_reporting_passthrough(upstream_headers, status_code=200): +async def _run_upstream_reporting_passthrough( + upstream_headers, status_code=200, cost_per_request=None +): """Drive a generic pass-through against an upstream that reports its own cost/usage. Returns (recorded standard logging payloads, proxy logging mock).""" from litellm.proxy._types import UserAPIKeyAuth @@ -4694,6 +4696,7 @@ async def _run_upstream_reporting_passthrough(upstream_headers, status_code=200) user_api_key_dict=UserAPIKeyAuth( api_key="sk-upstream-usage", team_id="team-fil" ), + cost_per_request=cost_per_request, ) for coroutine in enqueued: await coroutine @@ -4829,3 +4832,31 @@ async def test_streaming_passthrough_records_cost_and_tokens_reported_by_upstrea assert len(recorder.payloads) == 1 assert recorder.payloads[0]["response_cost"] == 0.00312 assert recorder.payloads[0]["total_tokens"] == 4021 + + +@pytest.mark.asyncio +async def test_upstream_reported_cost_survives_default_cost_per_request(): + """ + PassThroughGenericEndpoint.cost_per_request defaults to 0.0, so every + config-defined endpoint forwards a 0.0 flat cost even when the operator + never configured one. That flat estimate must not overwrite the real cost + the upstream reported for the request. + """ + payloads, _ = await _run_upstream_reporting_passthrough( + { + "x-litellm-response-cost": "0.000415", + "x-litellm-total-tokens": "1874", + }, + cost_per_request=0.0, + ) + + assert len(payloads) == 1 + assert payloads[0]["response_cost"] == 0.000415 + + +@pytest.mark.asyncio +async def test_configured_cost_per_request_still_applies_without_usage_headers(): + payloads, _ = await _run_upstream_reporting_passthrough({}, cost_per_request=0.25) + + assert len(payloads) == 1 + assert payloads[0]["response_cost"] == 0.25 From 9c48ad41ace598dadd6d93c0a9092d6f9531354e Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 24 Jul 2026 18:49:41 -0700 Subject: [PATCH 3/3] fix(passthrough): honor the zero fallback and aggregate-only TPM usage Two follow-ups from review on the upstream-reported usage contract. An unusable cost header fell through to the endpoint's flat cost_per_request instead of the zero the contract promises, so a target that contradicted itself got billed an estimate it had just disowned. A target that speaks this contract now owns the cost for the request whether or not the value it sent parsed. The reported total also cannot be split into prompt and completion, so reading one out of it under token_rate_limit_type input or output yielded zero and left the TPM window uncharged; pass-through traffic then ran past a limit it is meant to share with the general API. Usage that carries no split now charges its total under every limit type, while usage that does carry one is untouched. --- .../hooks/parallel_request_limiter_v3.py | 31 +++++++ .../pass_through_endpoints/success_handler.py | 10 ++- .../upstream_usage_headers.py | 14 +-- .../hooks/test_parallel_request_limiter_v3.py | 86 +++++++++++++++++++ .../test_pass_through_endpoints.py | 17 ++++ 5 files changed, 148 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 81d46247507..f72492881b3 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -2671,6 +2671,35 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return total_tokens + @staticmethod + def _aggregate_only_total_tokens(usage: Union[Usage, dict, None]) -> int: + """Total for usage that carries no input/output split, else 0. + + A source that can only report one number for the whole request (a + pass-through target pricing its own multi-model call) charges that + number under every ``token_rate_limit_type``. Splitting it is + impossible, and reading 0 out of it would leave the window + uncharged, which is how pass-through traffic slips past a TPM limit + it is supposed to share. + """ + if isinstance(usage, Usage): + prompt_tokens, completion_tokens, total_tokens = ( + usage.prompt_tokens or 0, + usage.completion_tokens or 0, + usage.total_tokens or 0, + ) + elif isinstance(usage, dict): + prompt_tokens, completion_tokens, total_tokens = ( + usage.get("prompt_tokens") or 0, + usage.get("completion_tokens") or 0, + usage.get("total_tokens") or 0, + ) + else: + return 0 + if prompt_tokens or completion_tokens: + return 0 + return total_tokens + async def _execute_token_increment_script( self, pipeline_operations: List["RedisPipelineIncrementOperation"], @@ -3107,6 +3136,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): if isinstance(_combined_usage, Usage): _usage = _combined_usage total_tokens = self._get_total_tokens_from_usage(usage=_usage, rate_limit_type=rate_limit_type) + if total_tokens == 0: + total_tokens = self._aggregate_only_total_tokens(usage=_usage) reserved_tokens = self._get_reserved_tokens_from_kwargs( kwargs=kwargs, diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 84a956f0c29..20ed84c8636 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -30,7 +30,7 @@ from .llm_provider_handlers.gemini_passthrough_logging_handler import ( from .llm_provider_handlers.vertex_passthrough_logging_handler import ( VertexPassthroughLoggingHandler, ) -from .upstream_usage_headers import upstream_reported_cost +from .upstream_usage_headers import has_upstream_reported_usage cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler() @@ -450,15 +450,17 @@ class PassThroughEndpointLogging: Only set the cost per request if it's set in the passthrough logging payload. If it's not set, don't set it in the logging object. - An upstream that priced the request itself always wins: ``cost_per_request`` + An upstream that prices its own requests always wins: ``cost_per_request`` is a flat per-request estimate for targets LiteLLM cannot price, and it defaults to 0.0 on every config-defined endpoint, so honoring it here - would zero out the real cost the upstream reported. + would zero out the real cost the upstream reported. That holds even when + the reported value was unusable, where the contract records 0 rather + than billing an estimate the upstream just contradicted. """ ######################################################### # Check if cost per request is set ######################################################### - if upstream_reported_cost(logging_obj) is not None: + if has_upstream_reported_usage(logging_obj): return kwargs if passthrough_logging_payload.get("cost_per_request") is not None: diff --git a/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py index e147dcb9547..6ee2514b4a0 100644 --- a/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py +++ b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py @@ -126,9 +126,11 @@ def apply_upstream_reported_usage( return reported -def upstream_reported_cost(logging_obj: LiteLLMLoggingObj) -> float | None: - """The cost the upstream reported for this request, if it reported one.""" - reported = logging_obj.model_call_details.get(UPSTREAM_REPORTED_USAGE_KEY) - if not isinstance(reported, UpstreamReportedUsage): - return None - return reported.response_cost +def has_upstream_reported_usage(logging_obj: LiteLLMLoggingObj) -> bool: + """Whether the upstream spoke this contract on the request's response. + + True even when the value it sent was unusable: a target that reports its + own totals owns the cost for the request, and a header we could not parse + means zero, never a fallback to someone else's estimate. + """ + return isinstance(logging_obj.model_call_details.get(UPSTREAM_REPORTED_USAGE_KEY), UpstreamReportedUsage) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 27b1c025fbb..9337050b61c 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -4702,3 +4702,89 @@ async def test_async_log_success_event_counts_passthrough_reported_tokens(monkey assert any("team-fil" in op["key"] for op in token_operations), ( "team TPM window must be charged so pass-through shares the team's limit" ) + + +@pytest.mark.parametrize("rate_limit_type", ["input", "output", "total"]) +@pytest.mark.asyncio +async def test_aggregate_only_usage_charges_tpm_under_every_limit_type( + monkeypatch, rate_limit_type +): + """ + A pass-through target reports one total for the whole request and cannot + split it into prompt/completion. Reading a split out of it yields 0, which + left the TPM window uncharged under input- or output-token limiting and let + pass-through traffic run past a limit it is supposed to share. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + + _api_key = hash_token("sk-aggregate-only") + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + monkeypatch.setattr( + parallel_request_handler, "get_rate_limit_type", lambda: rate_limit_type + ) + + captured_operations = [] + + async def mock_increment_pipeline(increment_list, **kwargs): + captured_operations.extend(increment_list) + return True + + monkeypatch.setattr( + parallel_request_handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + mock_increment_pipeline, + ) + + await parallel_request_handler.async_log_success_event( + kwargs={ + "standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}, + "combined_usage_object": Usage(total_tokens=1874), + }, + response_obj={"response": "upstream body was never parsed"}, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + token_operations = [op for op in captured_operations if op["key"].endswith(":tokens")] + assert token_operations, f"aggregate usage should charge TPM under {rate_limit_type}" + assert all(op["increment_value"] == 1874 for op in token_operations) + + +@pytest.mark.asyncio +async def test_split_usage_still_respects_the_configured_limit_type(monkeypatch): + """The aggregate fallback must not hijack usage that does carry a split.""" + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + + _api_key = hash_token("sk-split-usage") + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + monkeypatch.setattr(parallel_request_handler, "get_rate_limit_type", lambda: "output") + + captured_operations = [] + + async def mock_increment_pipeline(increment_list, **kwargs): + captured_operations.extend(increment_list) + return True + + monkeypatch.setattr( + parallel_request_handler.internal_usage_cache.dual_cache, + "async_increment_cache_pipeline", + mock_increment_pipeline, + ) + + await parallel_request_handler.async_log_success_event( + kwargs={"standard_logging_object": {"metadata": {"user_api_key_hash": _api_key}}}, + response_obj=ModelResponse( + model="gpt-4o", + usage=Usage(prompt_tokens=100, completion_tokens=7, total_tokens=107), + ), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + token_operations = [op for op in captured_operations if op["key"].endswith(":tokens")] + assert token_operations + assert all(op["increment_value"] == 7 for op in token_operations) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 60340cfa47b..9bddeda0723 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -4860,3 +4860,20 @@ async def test_configured_cost_per_request_still_applies_without_usage_headers() assert len(payloads) == 1 assert payloads[0]["response_cost"] == 0.25 + + +@pytest.mark.asyncio +async def test_unusable_upstream_cost_records_zero_not_the_flat_estimate(): + """ + A target that speaks this contract owns the cost for the request. When the + value it sends is unusable the contract records 0, rather than falling back + to a flat cost_per_request the upstream just contradicted. + """ + payloads, _ = await _run_upstream_reporting_passthrough( + {"x-litellm-response-cost": "not-a-number", "x-litellm-total-tokens": "1874"}, + cost_per_request=0.05, + ) + + assert len(payloads) == 1 + assert payloads[0]["response_cost"] == 0.0 + assert payloads[0]["total_tokens"] == 1874