Merge pull request #34590 from BerriAI/litellm_passthrough_upstream_reported_usage

feat(passthrough): record cost and usage reported by the upstream target
This commit is contained in:
yuneng-jiang 2026-07-25 09:06:09 -07:00 • committed by GitHub
commit c63fb4cb42
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 751 additions and 5 deletions

View file

@ -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)

View file

@ -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"],
@ -3086,8 +3115,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 +3131,13 @@ 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)
if total_tokens == 0:
total_tokens = self._aggregate_only_total_tokens(usage=_usage)
reserved_tokens = self._get_reserved_tokens_from_kwargs(
kwargs=kwargs,

View file

@ -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(

View file

@ -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 has_upstream_reported_usage
cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler()
@ -448,10 +449,20 @@ 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 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. 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 has_upstream_reported_usage(logging_obj):
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")

View file

@ -0,0 +1,136 @@
"""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"
# 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:
"""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
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 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)

View file

@ -4652,3 +4652,139 @@ 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"
)
@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)

View file

@ -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,262 @@ 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, 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
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"
),
cost_per_request=cost_per_request,
)
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
@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
@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

View file

@ -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)