fix(proxy): return cost breakdown header values as a named tuple

This commit is contained in:
mateo-berri 2026-08-14 17:04:26 -07:00
parent 1241bd5ce1
commit 9079e4c47b
2 changed files with 77 additions and 100 deletions

View file

@ -8,7 +8,7 @@ from collections.abc import AsyncGenerator, Callable, Mapping
from datetime import datetime
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, overload
import anyio
import httpx
@ -928,50 +928,41 @@ def _override_openai_response_model(
)
class CostBreakdownHeaderValues(NamedTuple):
original_cost: float | None = None
discount_amount: float | None = None
margin_total_amount: float | None = None
margin_percent: float | None = None
input_cost: float | None = None
output_cost: float | None = None
cache_read_cost: float | None = None
cache_creation_cost: float | None = None
reasoning_cost: float | None = None
tool_usage_cost: float | None = None
def _get_cost_breakdown_from_logging_obj(
litellm_logging_obj: LiteLLMLoggingObj | None,
) -> tuple[
float | None,
float | None,
float | None,
float | None,
float | None,
float | None,
float | None,
float | None,
float | None,
float | None,
]:
) -> CostBreakdownHeaderValues:
"""Extract discount, margin, and per-component cost information from logging object's cost breakdown."""
if not litellm_logging_obj or not hasattr(litellm_logging_obj, "cost_breakdown"):
return None, None, None, None, None, None, None, None, None, None
return CostBreakdownHeaderValues()
cost_breakdown: Final = litellm_logging_obj.cost_breakdown
if not cost_breakdown:
return None, None, None, None, None, None, None, None, None, None
return CostBreakdownHeaderValues()
original_cost: Final = cost_breakdown.get("original_cost")
discount_amount: Final = cost_breakdown.get("discount_amount")
margin_total_amount: Final = cost_breakdown.get("margin_total_amount")
margin_percent: Final = cost_breakdown.get("margin_percent")
input_cost: Final = cost_breakdown.get("input_cost")
output_cost: Final = cost_breakdown.get("output_cost")
cache_read_cost: Final = cost_breakdown.get("cache_read_cost")
cache_creation_cost: Final = cost_breakdown.get("cache_creation_cost")
reasoning_cost: Final = cost_breakdown.get("reasoning_cost")
tool_usage_cost: Final = cost_breakdown.get("tool_usage_cost")
return (
original_cost,
discount_amount,
margin_total_amount,
margin_percent,
input_cost,
output_cost,
cache_read_cost,
cache_creation_cost,
reasoning_cost,
tool_usage_cost,
return CostBreakdownHeaderValues(
original_cost=cost_breakdown.get("original_cost"),
discount_amount=cost_breakdown.get("discount_amount"),
margin_total_amount=cost_breakdown.get("margin_total_amount"),
margin_percent=cost_breakdown.get("margin_percent"),
input_cost=cost_breakdown.get("input_cost"),
output_cost=cost_breakdown.get("output_cost"),
cache_read_cost=cost_breakdown.get("cache_read_cost"),
cache_creation_cost=cost_breakdown.get("cache_creation_cost"),
reasoning_cost=cost_breakdown.get("reasoning_cost"),
tool_usage_cost=cost_breakdown.get("tool_usage_cost"),
)
@ -1098,19 +1089,7 @@ class ProxyBaseLLMRequestProcessing:
exclude_values: Final = {"", None, "None"}
hidden_params = hidden_params or {}
# Extract discount, margin, and per-component cost info from cost_breakdown if available
(
original_cost,
discount_amount,
margin_total_amount,
margin_percent,
input_cost,
output_cost,
cache_read_cost,
cache_creation_cost,
reasoning_cost,
tool_usage_cost,
) = _get_cost_breakdown_from_logging_obj(litellm_logging_obj=litellm_logging_obj)
cost_breakdown: Final = _get_cost_breakdown_from_logging_obj(litellm_logging_obj=litellm_logging_obj)
# Calculate updated spend for header (include current response_cost)
current_spend: Final = user_api_key_dict.spend or 0.0
@ -1139,20 +1118,36 @@ class ProxyBaseLLMRequestProcessing:
"x-litellm-version": version,
"x-litellm-model-region": model_region,
"x-litellm-response-cost": str(response_cost),
"x-litellm-response-cost-original": (str(original_cost) if original_cost is not None else None),
"x-litellm-response-cost-discount-amount": (str(discount_amount) if discount_amount is not None else None),
"x-litellm-response-cost-original": (
str(cost_breakdown.original_cost) if cost_breakdown.original_cost is not None else None
),
"x-litellm-response-cost-discount-amount": (
str(cost_breakdown.discount_amount) if cost_breakdown.discount_amount is not None else None
),
"x-litellm-response-cost-margin-amount": (
str(margin_total_amount) if margin_total_amount is not None else None
str(cost_breakdown.margin_total_amount) if cost_breakdown.margin_total_amount is not None else None
),
"x-litellm-response-cost-margin-percent": (
str(cost_breakdown.margin_percent) if cost_breakdown.margin_percent is not None else None
),
"x-litellm-response-cost-input": (
str(cost_breakdown.input_cost) if cost_breakdown.input_cost is not None else None
),
"x-litellm-response-cost-output": (
str(cost_breakdown.output_cost) if cost_breakdown.output_cost is not None else None
),
"x-litellm-response-cost-cache-read": (
str(cost_breakdown.cache_read_cost) if cost_breakdown.cache_read_cost is not None else None
),
"x-litellm-response-cost-margin-percent": (str(margin_percent) if margin_percent is not None else None),
"x-litellm-response-cost-input": (str(input_cost) if input_cost is not None else None),
"x-litellm-response-cost-output": (str(output_cost) if output_cost is not None else None),
"x-litellm-response-cost-cache-read": (str(cache_read_cost) if cache_read_cost is not None else None),
"x-litellm-response-cost-cache-creation": (
str(cache_creation_cost) if cache_creation_cost is not None else None
str(cost_breakdown.cache_creation_cost) if cost_breakdown.cache_creation_cost is not None else None
),
"x-litellm-response-cost-reasoning": (
str(cost_breakdown.reasoning_cost) if cost_breakdown.reasoning_cost is not None else None
),
"x-litellm-response-cost-tool-usage": (
str(cost_breakdown.tool_usage_cost) if cost_breakdown.tool_usage_cost is not None else None
),
"x-litellm-response-cost-reasoning": (str(reasoning_cost) if reasoning_cost is not None else None),
"x-litellm-response-cost-tool-usage": (str(tool_usage_cost) if tool_usage_cost is not None else None),
"x-litellm-classifier-cost": (str(classifier_cost) if classifier_cost is not None else None),
"x-litellm-key-tpm-limit": str(user_api_key_dict.tpm_limit),
"x-litellm-key-rpm-limit": str(user_api_key_dict.rpm_limit),

View file

@ -1087,16 +1087,14 @@ class TestProxyBaseLLMRequestProcessing:
discount_amount=0.000005,
)
(
original_cost,
discount_amount,
margin_total_amount,
margin_percent,
) = _get_cost_breakdown_from_logging_obj(logging_obj)
assert original_cost == 0.0001
assert discount_amount == 0.000005
assert margin_total_amount is None
assert margin_percent is None
breakdown = _get_cost_breakdown_from_logging_obj(logging_obj)
assert breakdown.original_cost == 0.0001
assert breakdown.discount_amount == 0.000005
assert breakdown.margin_total_amount is None
assert breakdown.margin_percent is None
assert breakdown.input_cost == 0.00005
assert breakdown.output_cost == 0.00005
assert breakdown.tool_usage_cost == 0.0
# Test with margin info
logging_obj_with_margin = LiteLLMLoggingObj(
@ -1118,16 +1116,11 @@ class TestProxyBaseLLMRequestProcessing:
margin_total_amount=0.00001,
)
(
original_cost,
discount_amount,
margin_total_amount,
margin_percent,
) = _get_cost_breakdown_from_logging_obj(logging_obj_with_margin)
assert original_cost == 0.0001
assert discount_amount is None
assert margin_total_amount == 0.00001
assert margin_percent == 0.10
breakdown_with_margin = _get_cost_breakdown_from_logging_obj(logging_obj_with_margin)
assert breakdown_with_margin.original_cost == 0.0001
assert breakdown_with_margin.discount_amount is None
assert breakdown_with_margin.margin_total_amount == 0.00001
assert breakdown_with_margin.margin_percent == 0.10
# Test with no discount or margin info
logging_obj_no_discount = LiteLLMLoggingObj(
@ -1146,28 +1139,17 @@ class TestProxyBaseLLMRequestProcessing:
cost_for_built_in_tools_cost_usd_dollar=0.0,
)
(
original_cost,
discount_amount,
margin_total_amount,
margin_percent,
) = _get_cost_breakdown_from_logging_obj(logging_obj_no_discount)
assert original_cost is None
assert discount_amount is None
assert margin_total_amount is None
assert margin_percent is None
breakdown_no_discount = _get_cost_breakdown_from_logging_obj(logging_obj_no_discount)
assert breakdown_no_discount.original_cost is None
assert breakdown_no_discount.discount_amount is None
assert breakdown_no_discount.margin_total_amount is None
assert breakdown_no_discount.margin_percent is None
assert breakdown_no_discount.input_cost == 0.00005
assert breakdown_no_discount.output_cost == 0.00005
# Test with None logging object
(
original_cost,
discount_amount,
margin_total_amount,
margin_percent,
) = _get_cost_breakdown_from_logging_obj(None)
assert original_cost is None
assert discount_amount is None
assert margin_total_amount is None
assert margin_percent is None
breakdown_none = _get_cost_breakdown_from_logging_obj(None)
assert all(value is None for value in breakdown_none)
def test_get_custom_headers_key_spend_includes_response_cost(self):
"""