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..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"], @@ -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, 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/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 932216141fe..20ed84c8636 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 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") 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..6ee2514b4a0 --- /dev/null +++ b/litellm/proxy/pass_through_endpoints/upstream_usage_headers.py @@ -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) 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..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 @@ -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) 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..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 @@ -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 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)