fix: expose cost breakdown in paid responses

This commit is contained in:
9qeklajc
2026-07-24 22:24:24 +02:00
parent b94d95fc2b
commit 81c0ff57e9
5 changed files with 175 additions and 14 deletions
+28 -2
View File
@@ -481,12 +481,38 @@ def _calculate_from_usd_cost(
) )
output_msats = cost_in_msats - input_msats output_msats = cost_in_msats - input_msats
# Estimate cache read/creation msats proportionally within the input cost.
# These are informational subcomponents: input_msats remains inclusive of
# cache cost so input_msats + output_msats == total_msats, matching the
# token-priced path and the public CostData contract.
cache_read_msats = 0
cache_creation_msats = 0
if cache_read_tokens > 0 or cache_creation_tokens > 0:
cache_tokens = cache_read_tokens + cache_creation_tokens
regular_input_tokens = input_tokens
total_input_tokens = regular_input_tokens + cache_tokens
if total_input_tokens > 0:
# Approximate by token count because the USD path only exposes an
# aggregate input cost, not separately priced cache buckets.
cache_read_msats = (
int(input_msats * cache_read_tokens / total_input_tokens)
if cache_read_tokens > 0
else 0
)
cache_creation_msats = (
int(input_msats * cache_creation_tokens / total_input_tokens)
if cache_creation_tokens > 0
else 0
)
logger.info( logger.info(
"Using cost from usage data/details", "Using cost from usage data/details",
extra={ extra={
"usd_cost": usd_cost, "usd_cost": usd_cost,
"cost_in_sats": cost_in_sats, "cost_in_sats": cost_in_sats,
"cost_in_msats": cost_in_msats, "cost_in_msats": cost_in_msats,
"cache_read_msats": cache_read_msats,
"cache_creation_msats": cache_creation_msats,
"model": response_data.get("model", "unknown"), "model": response_data.get("model", "unknown"),
}, },
) )
@@ -501,8 +527,8 @@ def _calculate_from_usd_cost(
output_tokens=output_tokens, output_tokens=output_tokens,
cache_read_input_tokens=cache_read_tokens, cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens=cache_creation_tokens, cache_creation_input_tokens=cache_creation_tokens,
cache_read_msats=0, cache_read_msats=cache_read_msats,
cache_creation_msats=0, cache_creation_msats=cache_creation_msats,
) )
+132 -8
View File
@@ -69,6 +69,56 @@ if typing.TYPE_CHECKING:
logger = get_logger(__name__) logger = get_logger(__name__)
def _inject_cost_response_headers(
headers: dict[str, str], cost_data: CostData | MaxCostData
) -> None:
"""Inject per-request cost breakdown into response headers.
The SDK's ``extractUsageFromResponseHeaders`` reads these to populate
``inputMsats``, ``outputMsats``, ``totalMsats`` and ``satsCost`` in the
usage tracking entry — without them, x-cashu requests show 0.0 for all
sat cost fields.
"""
headers["X-Routstr-Cost-Msats"] = str(cost_data.total_msats)
headers["X-Routstr-Input-Cost-Msats"] = str(cost_data.input_msats)
headers["X-Routstr-Output-Cost-Msats"] = str(cost_data.output_msats)
if cost_data.total_usd:
headers["X-Routstr-Cost-Usd"] = str(cost_data.total_usd)
def _inject_cost_into_usage(
response_json: dict, cost_data: CostData | MaxCostData
) -> None:
"""Inject cost breakdown into the response body's ``usage.cost`` object.
The SDK's ``extractUsageFromResponseBody`` expects ``usage.cost`` to be
an object with ``total_msats``/``input_msats``/``output_msats`` (not a
plain USD number). When the upstream returns ``cost`` as a number, the
SDK cannot extract the msats breakdown from the body alone.
"""
usage = response_json.get("usage")
if not isinstance(usage, dict):
return
# Direct assignment (not setdefault) so routstr's authoritative cost
# data always overwrites any upstream-provided cost values. Using
# setdefault would silently keep stale upstream values and drop our
# calculated msats breakdown.
cost_obj: dict[str, int | float] = {
"base_msats": cost_data.base_msats,
"input_msats": cost_data.input_msats,
"output_msats": cost_data.output_msats,
"total_msats": cost_data.total_msats,
"cache_read_input_tokens": cost_data.cache_read_input_tokens,
"cache_creation_input_tokens": cost_data.cache_creation_input_tokens,
"cache_read_msats": cost_data.cache_read_msats,
"cache_creation_msats": cost_data.cache_creation_msats,
}
if cost_data.total_usd:
cost_obj["total_usd"] = cost_data.total_usd
usage["cost"] = cost_obj
usage["cost_sats"] = cost_data.total_msats // 1000
def _is_json_content_type(content_type: str | None) -> bool: def _is_json_content_type(content_type: str | None) -> bool:
"""Return True when the upstream response should be parsed as JSON.""" """Return True when the upstream response should be parsed as JSON."""
if not content_type: if not content_type:
@@ -284,9 +334,27 @@ class BaseUpstreamProvider:
sats_cost = total_msats // 1000 sats_cost = total_msats // 1000
# Build the cost object that the SDK's extractUsageFromResponseBody
# and extractUsageFromSSEJson expect: an object with total_msats,
# input_msats, output_msats, cache_read_msats, cache_creation_msats,
# etc. Setting usage.cost to a plain float (total_usd) means the SDK
# cannot extract the msats breakdown — cache_read_msats and
# cache_creation_msats in particular are lost.
cost_obj = {
"base_msats": cost_dict.get("base_msats", 0),
"input_msats": cost_dict.get("input_msats", 0),
"output_msats": cost_dict.get("output_msats", 0),
"total_msats": total_msats,
"total_usd": total_usd,
"cache_read_input_tokens": cost_dict.get("cache_read_input_tokens", 0),
"cache_creation_input_tokens": cost_dict.get("cache_creation_input_tokens", 0),
"cache_read_msats": cost_dict.get("cache_read_msats", 0),
"cache_creation_msats": cost_dict.get("cache_creation_msats", 0),
}
# Inject into top-level usage block (OpenAI/Anthropic style) # Inject into top-level usage block (OpenAI/Anthropic style)
if "usage" in response_json: if "usage" in response_json:
response_json["usage"]["cost"] = total_usd response_json["usage"]["cost"] = cost_obj
response_json["usage"]["cost_sats"] = sats_cost response_json["usage"]["cost_sats"] = sats_cost
response_json["usage"]["remaining_balance_msats"] = key.balance response_json["usage"]["remaining_balance_msats"] = key.balance
self._fold_cache_into_input_tokens(response_json["usage"]) self._fold_cache_into_input_tokens(response_json["usage"])
@@ -2153,6 +2221,21 @@ class BaseUpstreamProvider:
if k.lower() in allowed_headers if k.lower() in allowed_headers
} }
# Inject cost breakdown headers so the SDK's
# extractUsageFromResponseHeaders can populate
# inputMsats/outputMsats/totalMsats for balance-mode requests.
if isinstance(cost_data, dict):
_cost_data_obj = CostData(
base_msats=cost_data.get("base_msats", 0),
input_msats=cost_data.get("input_msats", 0),
output_msats=cost_data.get("output_msats", 0),
total_msats=cost_data.get("total_msats", 0),
total_usd=cost_data.get("total_usd", 0.0),
)
else:
_cost_data_obj = cost_data
_inject_cost_response_headers(response_headers, _cost_data_obj)
return Response( return Response(
content=json.dumps(response_json).encode(), content=json.dumps(response_json).encode(),
status_code=response.status_code, status_code=response.status_code,
@@ -2242,9 +2325,24 @@ class BaseUpstreamProvider:
) )
self.inject_cost_metadata(response_json, cost_data, key) self.inject_cost_metadata(response_json, cost_data, key)
# Inject cost breakdown headers for balance-mode requests.
if isinstance(cost_data, dict):
_cost_data_obj = CostData(
base_msats=cost_data.get("base_msats", 0),
input_msats=cost_data.get("input_msats", 0),
output_msats=cost_data.get("output_msats", 0),
total_msats=cost_data.get("total_msats", 0),
total_usd=cost_data.get("total_usd", 0.0),
)
else:
_cost_data_obj = cost_data
response_headers: dict[str, str] = {}
_inject_cost_response_headers(response_headers, _cost_data_obj)
return Response( return Response(
content=json.dumps(response_json).encode(), content=json.dumps(response_json).encode(),
status_code=200, status_code=200,
headers=response_headers,
media_type="application/json", media_type="application/json",
) )
@@ -2295,11 +2393,12 @@ class BaseUpstreamProvider:
and "usage" in response_json and "usage" in response_json
and isinstance(response_json["usage"], dict) and isinstance(response_json["usage"], dict)
): ):
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 _inject_cost_into_usage(response_json, cost_data)
self._fold_cache_into_input_tokens(response_json["usage"]) self._fold_cache_into_input_tokens(response_json["usage"])
response_headers: dict[str, str] = {} response_headers: dict[str, str] = {}
if cost_data: if cost_data:
_inject_cost_response_headers(response_headers, cost_data)
refund_amount = messages_dispatch.compute_refund( refund_amount = messages_dispatch.compute_refund(
amount, unit, cost_data.total_msats amount, unit, cost_data.total_msats
) )
@@ -3700,6 +3799,11 @@ class BaseUpstreamProvider:
"model": model, "model": model,
}, },
) )
# Inject cost breakdown headers so the SDK's
# extractUsageFromResponseHeaders can populate
# inputMsats/outputMsats/totalMsats for x-cashu requests.
_inject_cost_response_headers(response_headers, cost_data)
except Exception as e: except Exception as e:
logger.error( logger.error(
"Error calculating cost for streaming response", "Error calculating cost for streaming response",
@@ -3722,8 +3826,12 @@ class BaseUpstreamProvider:
if "provider" not in data_json: if "provider" not in data_json:
self._apply_provider_field(data_json) self._apply_provider_field(data_json)
changed = True changed = True
if cost_data and "usage" in data_json and data_json["usage"]: if (
data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 cost_data
and "usage" in data_json
and data_json["usage"]
):
_inject_cost_into_usage(data_json, cost_data)
changed = True changed = True
if changed: if changed:
lines[i] = "data: " + json.dumps(data_json) lines[i] = "data: " + json.dumps(data_json)
@@ -3777,7 +3885,10 @@ class BaseUpstreamProvider:
) )
if cost_data and "usage" in response_json: if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 # Inject cost breakdown into both the response body (so the
# SDK's body extractor picks up the msats breakdown) and the
# response headers (so the SDK's header extractor works too).
_inject_cost_into_usage(response_json, cost_data)
if not cost_data: if not cost_data:
logger.error( logger.error(
@@ -3808,6 +3919,8 @@ class BaseUpstreamProvider:
if "content-encoding" in response_headers: if "content-encoding" in response_headers:
del response_headers["content-encoding"] del response_headers["content-encoding"]
_inject_cost_response_headers(response_headers, cost_data)
if unit == "msat": if unit == "msat":
refund_amount = amount - cost_data.total_msats refund_amount = amount - cost_data.total_msats
elif unit == "sat": elif unit == "sat":
@@ -4681,6 +4794,11 @@ class BaseUpstreamProvider:
"model": model, "model": model,
}, },
) )
# Inject cost breakdown headers so the SDK's
# extractUsageFromResponseHeaders can populate
# inputMsats/outputMsats/totalMsats for x-cashu requests.
_inject_cost_response_headers(response_headers, cost_data)
except Exception as e: except Exception as e:
logger.error( logger.error(
"Error calculating cost for streaming Responses API response", "Error calculating cost for streaming Responses API response",
@@ -4703,8 +4821,12 @@ class BaseUpstreamProvider:
if "provider" not in data_json: if "provider" not in data_json:
self._apply_provider_field(data_json) self._apply_provider_field(data_json)
changed = True changed = True
if cost_data and "usage" in data_json and data_json["usage"]: if (
data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 cost_data
and "usage" in data_json
and data_json["usage"]
):
_inject_cost_into_usage(data_json, cost_data)
changed = True changed = True
if changed: if changed:
lines[i] = "data: " + json.dumps(data_json) lines[i] = "data: " + json.dumps(data_json)
@@ -4747,7 +4869,7 @@ class BaseUpstreamProvider:
) )
if cost_data and "usage" in response_json: if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000 _inject_cost_into_usage(response_json, cost_data)
if not cost_data: if not cost_data:
logger.error( logger.error(
@@ -4778,6 +4900,8 @@ class BaseUpstreamProvider:
if "content-encoding" in response_headers: if "content-encoding" in response_headers:
del response_headers["content-encoding"] del response_headers["content-encoding"]
_inject_cost_response_headers(response_headers, cost_data)
if unit == "msat": if unit == "msat":
refund_amount = amount - cost_data.total_msats refund_amount = amount - cost_data.total_msats
elif unit == "sat": elif unit == "sat":
@@ -529,6 +529,8 @@ async def test_openrouter_upstream_inference_cost_components_are_used() -> None:
assert isinstance(result, CostData) assert isinstance(result, CostData)
assert result.input_msats == 994 assert result.input_msats == 994
assert result.output_msats == 3477 assert result.output_msats == 3477
assert result.cache_read_msats == 758
assert result.cache_creation_msats == 0
assert result.input_msats + result.output_msats == result.total_msats == 4471 assert result.input_msats + result.output_msats == result.total_msats == 4471
@@ -574,6 +576,8 @@ async def test_ppq_byok_bills_upstream_inference_cost_plus_fee() -> None:
# Token normalisation (OpenAI dialect: cached included in prompt_tokens) # Token normalisation (OpenAI dialect: cached included in prompt_tokens)
assert result.input_tokens == 5070 # 164371 - 159301 assert result.input_tokens == 5070 # 164371 - 159301
assert result.cache_read_input_tokens == 159301 assert result.cache_read_input_tokens == 159301
assert result.cache_read_msats == 897966
assert result.cache_creation_msats == 0
assert result.output_tokens == 99 assert result.output_tokens == 99
+2 -1
View File
@@ -451,7 +451,8 @@ async def test_non_streaming_dispatches_via_litellm_and_returns_anthropic_respon
assert payload["model"] == "openai/gpt-4o-mini" # mapped back to requested assert payload["model"] == "openai/gpt-4o-mini" # mapped back to requested
assert payload["usage"]["input_tokens"] == 5 assert payload["usage"]["input_tokens"] == 5
assert payload["usage"]["output_tokens"] == 3 assert payload["usage"]["output_tokens"] == 3
assert payload["usage"]["cost"] == 0.0001 assert payload["usage"]["cost"]["total_msats"] == 1234
assert payload["usage"]["cost"]["total_usd"] == 0.0001
assert payload["usage"]["cost_sats"] == 1 assert payload["usage"]["cost_sats"] == 1
+9 -3
View File
@@ -67,8 +67,13 @@ async def test_non_streaming_includes_cost_sats() -> None:
) )
body = json.loads(response.body) body = json.loads(response.body)
assert "cost_sats" in body["usage"]
assert body["usage"]["cost_sats"] == 5 # 5000 msats // 1000 assert body["usage"]["cost_sats"] == 5 # 5000 msats // 1000
assert body["usage"]["cost"]["total_msats"] == 5000
assert body["usage"]["cost"]["input_msats"] == 3000
assert body["usage"]["cost"]["output_msats"] == 2000
assert response.headers["x-routstr-cost-msats"] == "5000"
assert response.headers["x-routstr-input-cost-msats"] == "3000"
assert response.headers["x-routstr-output-cost-msats"] == "2000"
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -96,7 +101,7 @@ async def test_non_streaming_cost_sats_value_rounds_down() -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_non_streaming_preserves_existing_usage_fields() -> None: async def test_non_streaming_preserves_tokens_and_replaces_upstream_cost() -> None:
provider = _make_provider() provider = _make_provider()
cost_data = _make_cost_data(total_msats=3000) cost_data = _make_cost_data(total_msats=3000)
@@ -127,7 +132,8 @@ async def test_non_streaming_preserves_existing_usage_fields() -> None:
assert usage["prompt_tokens"] == 100 assert usage["prompt_tokens"] == 100
assert usage["completion_tokens"] == 50 assert usage["completion_tokens"] == 50
assert usage["total_tokens"] == 150 assert usage["total_tokens"] == 150
assert usage["cost"] == 0.00015 assert usage["cost"]["total_msats"] == 3000
assert usage["cost"]["total_usd"] == 0.00025
assert usage["cost_sats"] == 3 assert usage["cost_sats"] == 3