mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
c63fb4cb42
8 changed files with 751 additions and 5 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
136
litellm/proxy/pass_through_endpoints/upstream_usage_headers.py
Normal file
136
litellm/proxy/pass_through_endpoints/upstream_usage_headers.py
Normal 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)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue